#![allow(non_snake_case)]
use core::ops::BitAnd;
use teeny_macros::kernel;
use teeny_triton::triton::{
types::{AddOffsets, Comparison, Tensor},
*,
};
#[kernel]
pub fn conv2d_bn_silu_gemm_forward<
T: Triton,
const BLOCK_M: i32,
const BLOCK_N: i32,
const BLOCK_K: i32,
const GROUP_M: i32,
>(
x_ptr: T::Pointer<f32>,
w_ptr: T::Pointer<f32>,
bn_scale_ptr: T::Pointer<f32>,
bn_shift_ptr: T::Pointer<f32>,
y_ptr: T::Pointer<f32>,
B: i32,
C_IN: i32,
C_OUT: i32,
M: i32, y_row_stride: i32,
) where
T::I32Tensor: Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::BoolTensor: BitAnd<Output = T::BoolTensor>,
T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
{
let pid = T::program_id(Axis::X);
let num_pid_m = T::cdiv(M, BLOCK_M);
let num_pid_n = T::cdiv(C_OUT, BLOCK_N);
let pids_per_batch = num_pid_m * num_pid_n;
let b = pid / pids_per_batch;
let pid_local = pid % pids_per_batch;
let num_pid_in_group = GROUP_M * num_pid_n;
let group_id = pid_local / num_pid_in_group;
let first_pid_m = group_id * GROUP_M;
let remaining_m = num_pid_m - first_pid_m;
let group_size_m = if remaining_m < GROUP_M {
remaining_m
} else {
GROUP_M
};
let pid_in_group = pid_local % num_pid_in_group;
let pid_m = first_pid_m + (pid_in_group % group_size_m);
let pid_n = pid_in_group / group_size_m;
let x_desc = T::make_tensor_descriptor(
x_ptr,
&[B * C_IN, M],
&[M, 1],
&[BLOCK_K, BLOCK_M],
Some(PaddingOption::Zero),
);
let w_desc = T::make_tensor_descriptor(
w_ptr,
&[C_OUT, C_IN],
&[C_IN, 1],
&[BLOCK_N, BLOCK_K],
Some(PaddingOption::Zero),
);
let mut acc = T::zeros::<f32>(&[BLOCK_N, BLOCK_M]);
let k_tiles = T::cdiv(C_IN, BLOCK_K);
for k in 0..k_tiles {
let x_tile = T::load_tensor_descriptor(x_desc, &[b * C_IN + k * BLOCK_K, pid_m * BLOCK_M]);
let w_tile = T::load_tensor_descriptor(w_desc, &[pid_n * BLOCK_N, k * BLOCK_K]);
acc = T::dot::<f32, f32>(w_tile, x_tile, Some(acc), InputPrecision::TF32, None);
}
let bn_off = T::arange(0, BLOCK_N) + pid_n * BLOCK_N;
let bn_n_mask = bn_off.lt(C_OUT);
let bn_scale = T::load(
bn_scale_ptr.add_offsets(bn_off),
Some(bn_n_mask),
Some(T::zeros::<f32>(&[BLOCK_N])),
&[],
None,
None,
None,
false,
);
let bn_shift = T::load(
bn_shift_ptr.add_offsets(bn_off),
Some(bn_n_mask),
Some(T::zeros::<f32>(&[BLOCK_N])),
&[],
None,
None,
None,
false,
);
let scale_2d = T::broadcast_to(T::expand_dims(bn_scale, 1), &[BLOCK_N, BLOCK_M]);
let shift_2d = T::broadcast_to(T::expand_dims(bn_shift, 1), &[BLOCK_N, BLOCK_M]);
let bn_out = scale_2d * acc + shift_2d;
let y = bn_out * T::sigmoid(bn_out);
let y_desc = T::make_tensor_descriptor(
y_ptr,
&[B * C_OUT, y_row_stride],
&[y_row_stride, 1],
&[BLOCK_N, BLOCK_M],
Some(PaddingOption::Zero),
);
T::store_tensor_descriptor(y_desc, &[b * C_OUT + pid_n * BLOCK_N, pid_m * BLOCK_M], y);
}
impl teeny_core::model::RuntimeOp for Conv2dBnSiluGemmForward {
fn n_activation_inputs(&self) -> usize {
1
}
fn param_shapes(&self, input_shapes: &[&[usize]], output_shape: &[usize]) -> Vec<Vec<usize>> {
let c_in = input_shapes[0][1];
let c_out = output_shape[1];
vec![vec![c_out, c_in], vec![c_out], vec![c_out]]
}
fn param_names(&self) -> &'static [&'static str] {
&["weight", "bn_scale", "bn_shift"]
}
fn forward_output_row_elems(&self, output_shape: &[usize]) -> usize {
output_shape[2] * output_shape[3]
}
fn forward_output_row_stride(&self, output_shape: &[usize]) -> usize {
let m = self.forward_output_row_elems(output_shape);
m.next_multiple_of(self.block_m as usize)
}
fn pack_args(
&self,
inputs: &[(teeny_core::model::RawPtr, &[usize])],
params: &[teeny_core::model::RawPtr],
output: teeny_core::model::RawPtr,
output_shape: &[usize],
output_row_stride: i32,
visitor: &mut dyn teeny_core::device::program::ArgVisitor,
) {
let input_shape = inputs[0].1;
let b = input_shape[0] as i32;
let c_in = input_shape[1] as i32;
let c_out = output_shape[1] as i32;
let m = (output_shape[2] * output_shape[3]) as i32; visitor.visit_ptr(inputs[0].0); visitor.visit_ptr(params[0]); visitor.visit_ptr(params[1]); visitor.visit_ptr(params[2]); visitor.visit_ptr(output); visitor.visit_i32(b);
visitor.visit_i32(c_in);
visitor.visit_i32(c_out);
visitor.visit_i32(m);
visitor.visit_i32(output_row_stride); }
fn block(&self) -> [u32; 3] {
[128, 1, 1]
}
fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
let b = output_shape[0];
let m = output_shape[2] * output_shape[3]; let c_out = output_shape[1];
let pm = m.div_ceil(self.block_m as usize);
let pn = c_out.div_ceil(self.block_n as usize);
[(b * pm * pn) as u32, 1, 1]
}
}