#![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_tiled_forward<
T: Triton,
const KH: i32,
const KW: i32,
const STRIDE_H: i32,
const STRIDE_W: i32,
const PAD_H: i32,
const PAD_W: i32,
const BLOCK_OW: i32,
const BLOCK_N: 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,
H: i32,
W: i32,
OH: i32,
OW: i32,
y_col_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_ow_tiles = T::cdiv(OW, BLOCK_OW);
let num_n_tiles = T::cdiv(C_OUT, BLOCK_N);
let ow_tile = pid % num_ow_tiles;
let tmp = pid / num_ow_tiles;
let n_tile = tmp % num_n_tiles;
let tmp2 = tmp / num_n_tiles;
let oh = tmp2 % OH;
let b = tmp2 / OH;
let ow_start = ow_tile * BLOCK_OW;
let c_out_start = n_tile * BLOCK_N;
let ow_range = T::arange(0, BLOCK_OW) + ow_start; let c_out_range = T::arange(0, BLOCK_N) + c_out_start;
let ow_mask = ow_range.lt(OW); let n_mask = c_out_range.lt(C_OUT);
let mut acc = T::zeros::<f32>(&[BLOCK_N, BLOCK_OW]);
let loop_bound = C_IN * KH * KW;
for idx in 0..loop_bound {
let kw = idx % KW;
let kh_cin = idx / KW;
let kh = kh_cin % KH;
let c_in_local = kh_cin / KH;
let ih = oh * STRIDE_H + kh - PAD_H;
let iw_range = ow_range * STRIDE_W + kw - PAD_W;
#[allow(clippy::erasing_op)]
let ih_t = ow_range * 0 + ih;
let h_in_bounds = ih_t.ge(0) & ih_t.lt(H);
let w_in_bounds = iw_range.ge(0) & iw_range.lt(W);
let x_load_mask = ow_mask & h_in_bounds & w_in_bounds;
let x_offsets = iw_range + ((b * C_IN + c_in_local) * H * W + ih * W);
let x_col = T::load(
x_ptr.add_offsets(x_offsets),
Some(x_load_mask),
Some(T::zeros::<f32>(&[BLOCK_OW])),
&[],
None,
None,
None,
false,
);
let k_scalar = (c_in_local * KH + kh) * KW + kw;
let w_offsets = c_out_range * (C_IN * KH * KW) + k_scalar;
let w_row = T::load(
w_ptr.add_offsets(w_offsets),
Some(n_mask),
Some(T::zeros::<f32>(&[BLOCK_N])),
&[],
None,
None,
None,
false,
);
let w_2d = T::broadcast_to(T::expand_dims(w_row, 1), &[BLOCK_N, BLOCK_OW]);
let x_2d = T::broadcast_to(T::expand_dims(x_col, 0), &[BLOCK_N, BLOCK_OW]);
acc = acc + w_2d * x_2d;
}
let bn_scale = T::load(
bn_scale_ptr.add_offsets(c_out_range),
Some(n_mask),
Some(T::zeros::<f32>(&[BLOCK_N])),
&[],
None,
None,
None,
false,
);
let bn_shift = T::load(
bn_shift_ptr.add_offsets(c_out_range),
Some(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_OW]);
let shift_2d = T::broadcast_to(T::expand_dims(bn_shift, 1), &[BLOCK_N, BLOCK_OW]);
let bn_out = scale_2d * acc + shift_2d;
let y = bn_out * T::sigmoid(bn_out);
let oh_ycs = OH * y_col_stride;
let y_desc = T::make_tensor_descriptor(
y_ptr,
&[B * C_OUT, oh_ycs],
&[oh_ycs, 1],
&[BLOCK_N, BLOCK_OW],
Some(PaddingOption::Zero),
);
T::store_tensor_descriptor(
y_desc,
&[b * C_OUT + c_out_start, oh * y_col_stride + ow_start],
y,
);
}
impl teeny_core::model::RuntimeOp for Conv2dBnSiluTiledForward {
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, self.kh as usize, self.kw as usize],
vec![c_out],
vec![c_out],
]
}
fn param_names(&self) -> &'static [&'static str] {
&["weight", "bn_scale", "bn_shift"]
}
fn forward_output_row_stride(&self, output_shape: &[usize]) -> usize {
let ow = output_shape[3];
ow.next_multiple_of(self.block_ow 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 oh = output_shape[2] as i32;
let y_col_stride = output_row_stride;
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(input_shape[0] as i32); visitor.visit_i32(input_shape[1] as i32); visitor.visit_i32(output_shape[1] as i32); visitor.visit_i32(input_shape[2] as i32); visitor.visit_i32(input_shape[3] as i32); visitor.visit_i32(oh); visitor.visit_i32(output_shape[3] as i32); visitor.visit_i32(y_col_stride); }
fn block(&self) -> [u32; 3] {
[128, 1, 1]
}
fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
let num_ow_tiles = output_shape[3].div_ceil(self.block_ow as usize);
let num_n_tiles = output_shape[1].div_ceil(self.block_n as usize);
[
(output_shape[0] * output_shape[2] * num_n_tiles * num_ow_tiles) as u32,
1,
1,
]
}
}