#![allow(non_snake_case)]
use teeny_core::dtype::{Float, Num};
use teeny_macros::kernel;
use teeny_triton::triton::{
types::{AddOffsets, Comparison},
*,
};
macro_rules! impl_reduce_num_runtime_op {
($Fwd:ident) => {
impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for $Fwd<D> {
fn n_activation_inputs(&self) -> usize {
1
}
fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
vec![]
}
fn pack_args(
&self,
inputs: &[(teeny_core::model::RawPtr, &[usize])],
_: &[teeny_core::model::RawPtr],
output: teeny_core::model::RawPtr,
output_shape: &[usize],
_: i32,
visitor: &mut dyn teeny_core::device::program::ArgVisitor,
) {
let n_outer: usize = output_shape.iter().product::<usize>().max(1);
let n_total: usize = inputs[0].1.iter().product();
let n_inner: usize = if n_outer > 0 {
n_total / n_outer
} else {
n_total
};
visitor.visit_ptr(inputs[0].0);
visitor.visit_ptr(output);
visitor.visit_i32(n_inner as i32);
visitor.visit_i32(n_outer as i32);
}
fn block(&self) -> [u32; 3] {
[self.block_inner as u32, 1, 1]
}
fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
let n_outer: usize = output_shape.iter().product::<usize>().max(1);
[n_outer as u32, 1, 1]
}
}
};
}
macro_rules! impl_reduce_float_runtime_op {
($Fwd:ident) => {
impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for $Fwd<D> {
fn n_activation_inputs(&self) -> usize {
1
}
fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
vec![]
}
fn pack_args(
&self,
inputs: &[(teeny_core::model::RawPtr, &[usize])],
_: &[teeny_core::model::RawPtr],
output: teeny_core::model::RawPtr,
output_shape: &[usize],
_: i32,
visitor: &mut dyn teeny_core::device::program::ArgVisitor,
) {
let n_outer: usize = output_shape.iter().product::<usize>().max(1);
let n_total: usize = inputs[0].1.iter().product();
let n_inner: usize = if n_outer > 0 {
n_total / n_outer
} else {
n_total
};
visitor.visit_ptr(inputs[0].0);
visitor.visit_ptr(output);
visitor.visit_i32(n_inner as i32);
visitor.visit_i32(n_outer as i32);
}
fn block(&self) -> [u32; 3] {
[self.block_inner as u32, 1, 1]
}
fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
let n_outer: usize = output_shape.iter().product::<usize>().max(1);
[n_outer as u32, 1, 1]
}
}
};
}
#[kernel]
pub fn reduce_sum_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_inner: i32,
n_outer: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let row = T::program_id(Axis::X);
if row >= n_outer {
return;
}
let col_offsets = T::arange(0, BLOCK_INNER);
let offsets = col_offsets + row * n_inner;
let mask = col_offsets.lt(n_inner);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(T::zeros::<D>(&[BLOCK_INNER])),
&[],
None,
None,
None,
false,
);
let sum = T::sum(x, Some(0), true); let row_offsets = T::arange(0, 1) + row;
T::store(y_ptr.add_offsets(row_offsets), sum, None, &[], None, None);
}
impl_reduce_num_runtime_op!(ReduceSumForward);
#[kernel]
pub fn reduce_mean_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_inner: i32,
n_outer: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let row = T::program_id(Axis::X);
if row >= n_outer {
return;
}
let col_offsets = T::arange(0, BLOCK_INNER);
let offsets = col_offsets + row * n_inner;
let mask = col_offsets.lt(n_inner);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(T::zeros::<D>(&[BLOCK_INNER])),
&[],
None,
None,
None,
false,
);
let sum = T::sum(x, Some(0), true);
let n_f = T::cast::<i32, D>(T::full::<i32>(&[1], n_inner), None, false);
let mean = sum / n_f;
let row_offsets = T::arange(0, 1) + row;
T::store(y_ptr.add_offsets(row_offsets), mean, None, &[], None, None);
}
impl_reduce_float_runtime_op!(ReduceMeanForward);
#[kernel]
pub fn reduce_max_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_inner: i32,
n_outer: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let row = T::program_id(Axis::X);
if row >= n_outer {
return;
}
let col_offsets = T::arange(0, BLOCK_INNER);
let offsets = col_offsets + row * n_inner;
let mask = col_offsets.lt(n_inner);
let neg_inf = T::cast::<f32, D>(
T::full::<f32>(&[BLOCK_INNER], -3.4028235e38_f32),
None,
false,
);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(neg_inf),
&[],
None,
None,
None,
false,
);
let val = T::max(x, Some(0), true);
let row_offsets = T::arange(0, 1) + row;
T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
}
impl_reduce_num_runtime_op!(ReduceMaxForward);
#[kernel]
pub fn reduce_min_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_inner: i32,
n_outer: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let row = T::program_id(Axis::X);
if row >= n_outer {
return;
}
let col_offsets = T::arange(0, BLOCK_INNER);
let offsets = col_offsets + row * n_inner;
let mask = col_offsets.lt(n_inner);
let pos_inf = T::cast::<f32, D>(
T::full::<f32>(&[BLOCK_INNER], 3.4028235e38_f32),
None,
false,
);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(pos_inf),
&[],
None,
None,
None,
false,
);
let val = T::min(x, Some(0), true);
let row_offsets = T::arange(0, 1) + row;
T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
}
impl_reduce_num_runtime_op!(ReduceMinForward);
#[kernel]
pub fn reduce_l1_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_inner: i32,
n_outer: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let row = T::program_id(Axis::X);
if row >= n_outer {
return;
}
let col_offsets = T::arange(0, BLOCK_INNER);
let offsets = col_offsets + row * n_inner;
let mask = col_offsets.lt(n_inner);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(T::zeros::<D>(&[BLOCK_INNER])),
&[],
None,
None,
None,
false,
);
let val = T::sum(T::abs(x), Some(0), true);
let row_offsets = T::arange(0, 1) + row;
T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
}
impl_reduce_num_runtime_op!(ReduceL1Forward);
#[kernel]
pub fn reduce_l2_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_inner: i32,
n_outer: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let row = T::program_id(Axis::X);
if row >= n_outer {
return;
}
let col_offsets = T::arange(0, BLOCK_INNER);
let offsets = col_offsets + row * n_inner;
let mask = col_offsets.lt(n_inner);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(T::zeros::<D>(&[BLOCK_INNER])),
&[],
None,
None,
None,
false,
);
let sum_sq = T::sum(x * x, Some(0), true);
let val = T::sqrt(sum_sq);
let row_offsets = T::arange(0, 1) + row;
T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
}
impl_reduce_float_runtime_op!(ReduceL2Forward);
#[kernel]
pub fn reduce_sum_square_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_inner: i32,
n_outer: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let row = T::program_id(Axis::X);
if row >= n_outer {
return;
}
let col_offsets = T::arange(0, BLOCK_INNER);
let offsets = col_offsets + row * n_inner;
let mask = col_offsets.lt(n_inner);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(T::zeros::<D>(&[BLOCK_INNER])),
&[],
None,
None,
None,
false,
);
let val = T::sum(x * x, Some(0), true);
let row_offsets = T::arange(0, 1) + row;
T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
}
impl_reduce_num_runtime_op!(ReduceSumSquareForward);
#[kernel]
pub fn reduce_log_sum_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_inner: i32,
n_outer: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let row = T::program_id(Axis::X);
if row >= n_outer {
return;
}
let col_offsets = T::arange(0, BLOCK_INNER);
let offsets = col_offsets + row * n_inner;
let mask = col_offsets.lt(n_inner);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(T::zeros::<D>(&[BLOCK_INNER])),
&[],
None,
None,
None,
false,
);
let sum = T::sum(x, Some(0), true);
let val = T::log(sum);
let row_offsets = T::arange(0, 1) + row;
T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
}
impl_reduce_float_runtime_op!(ReduceLogSumForward);
#[kernel]
pub fn reduce_log_sum_exp_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_inner: i32,
n_outer: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let row = T::program_id(Axis::X);
if row >= n_outer {
return;
}
let col_offsets = T::arange(0, BLOCK_INNER);
let offsets = col_offsets + row * n_inner;
let mask = col_offsets.lt(n_inner);
let neg_inf = T::cast::<f32, D>(
T::full::<f32>(&[BLOCK_INNER], -3.4028235e38_f32),
None,
false,
);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(neg_inf),
&[],
None,
None,
None,
false,
);
let m = T::max(x, Some(0), true); let fill = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_INNER], 0.0_f32), None, false);
let x_adj = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(fill),
&[],
None,
None,
None,
false,
);
let sum_exp = T::sum(T::exp(x_adj - m), Some(0), true);
let val = m + T::log(sum_exp);
let row_offsets = T::arange(0, 1) + row;
T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
}
impl_reduce_float_runtime_op!(ReduceLogSumExpForward);
#[kernel]
pub fn reduce_prod_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_inner: i32,
n_outer: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let row = T::program_id(Axis::X);
if row >= n_outer {
return;
}
let col_offsets = T::arange(0, BLOCK_INNER);
let offsets = col_offsets + row * n_inner;
let mask = col_offsets.lt(n_inner);
let one_fill = T::cast::<f32, D>(T::full::<f32>(&[BLOCK_INNER], 1.0_f32), None, false);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(one_fill),
&[],
None,
None,
None,
false,
);
let val = T::exp(T::sum(T::log(x), Some(0), true));
let row_offsets = T::arange(0, 1) + row;
T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
}
impl_reduce_float_runtime_op!(ReduceProdForward);
#[kernel]
pub fn cum_sum_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_inner: i32,
n_outer: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let row = T::program_id(Axis::X);
if row >= n_outer {
return;
}
let col_offsets = T::arange(0, BLOCK_INNER);
let offsets = col_offsets + row * n_inner;
let mask = col_offsets.lt(n_inner);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(T::zeros::<D>(&[BLOCK_INNER])),
&[],
None,
None,
None,
false,
);
let y = T::cumsum(x, 0, false);
T::store(y_ptr.add_offsets(offsets), y, Some(mask), &[], None, None);
}
impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for CumSumForward<D> {
fn n_activation_inputs(&self) -> usize {
1
}
fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
vec![]
}
fn pack_args(
&self,
inputs: &[(teeny_core::model::RawPtr, &[usize])],
_: &[teeny_core::model::RawPtr],
output: teeny_core::model::RawPtr,
output_shape: &[usize],
_: i32,
visitor: &mut dyn teeny_core::device::program::ArgVisitor,
) {
let n_total: usize = output_shape.iter().product();
let n_inner = output_shape.last().copied().unwrap_or(1);
let n_outer = n_total / n_inner;
visitor.visit_ptr(inputs[0].0);
visitor.visit_ptr(output);
visitor.visit_i32(n_inner as i32);
visitor.visit_i32(n_outer as i32);
}
fn block(&self) -> [u32; 3] {
[self.block_inner as u32, 1, 1]
}
fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
let n_total: usize = output_shape.iter().product();
let n_inner = output_shape.last().copied().unwrap_or(1);
let n_outer = n_total / n_inner;
[n_outer as u32, 1, 1]
}
}
#[kernel]
pub fn cum_prod_forward<T: Triton, D: Num, const BLOCK_INNER: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_inner: i32,
n_outer: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let row = T::program_id(Axis::X);
if row >= n_outer {
return;
}
let col_offsets = T::arange(0, BLOCK_INNER);
let offsets = col_offsets + row * n_inner;
let mask = col_offsets.lt(n_inner);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(T::zeros::<D>(&[BLOCK_INNER])),
&[],
None,
None,
None,
false,
);
let y = T::cumprod(x, 0, false);
T::store(y_ptr.add_offsets(offsets), y, Some(mask), &[], None, None);
}
impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for CumProdForward<D> {
fn n_activation_inputs(&self) -> usize {
1
}
fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
vec![]
}
fn pack_args(
&self,
inputs: &[(teeny_core::model::RawPtr, &[usize])],
_: &[teeny_core::model::RawPtr],
output: teeny_core::model::RawPtr,
output_shape: &[usize],
_: i32,
visitor: &mut dyn teeny_core::device::program::ArgVisitor,
) {
let n_total: usize = output_shape.iter().product();
let n_inner = output_shape.last().copied().unwrap_or(1);
let n_outer = n_total / n_inner;
visitor.visit_ptr(inputs[0].0);
visitor.visit_ptr(output);
visitor.visit_i32(n_inner as i32);
visitor.visit_i32(n_outer as i32);
}
fn block(&self) -> [u32; 3] {
[self.block_inner as u32, 1, 1]
}
fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
let n_total: usize = output_shape.iter().product();
let n_inner = output_shape.last().copied().unwrap_or(1);
let n_outer = n_total / n_inner;
[n_outer as u32, 1, 1]
}
}
#[kernel]
pub fn global_avg_pool_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_inner: i32,
n_outer: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let row = T::program_id(Axis::X);
if row >= n_outer {
return;
}
let col_offsets = T::arange(0, BLOCK_INNER);
let offsets = col_offsets + row * n_inner;
let mask = col_offsets.lt(n_inner);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(T::zeros::<D>(&[BLOCK_INNER])),
&[],
None,
None,
None,
false,
);
let sum = T::sum(x, Some(0), true);
let n_f = T::cast::<i32, D>(T::full::<i32>(&[1], n_inner), None, false);
let mean = sum / n_f;
let row_offsets = T::arange(0, 1) + row;
T::store(y_ptr.add_offsets(row_offsets), mean, None, &[], None, None);
}
impl_reduce_float_runtime_op!(GlobalAvgPoolForward);
#[kernel]
pub fn global_max_pool_forward<T: Triton, D: Float, const BLOCK_INNER: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_inner: i32,
n_outer: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
let row = T::program_id(Axis::X);
if row >= n_outer {
return;
}
let col_offsets = T::arange(0, BLOCK_INNER);
let offsets = col_offsets + row * n_inner;
let mask = col_offsets.lt(n_inner);
let neg_inf = T::cast::<f32, D>(
T::full::<f32>(&[BLOCK_INNER], -3.4028235e38_f32),
None,
false,
);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(mask),
Some(neg_inf),
&[],
None,
None,
None,
false,
);
let val = T::max(x, Some(0), true);
let row_offsets = T::arange(0, 1) + row;
T::store(y_ptr.add_offsets(row_offsets), val, None, &[], None, None);
}
impl_reduce_float_runtime_op!(GlobalMaxPoolForward);