use teeny_core::dtype::Num;
use teeny_macros::kernel;
use teeny_triton::triton::{Axis, PaddingOption, Triton};
#[kernel]
pub fn flatten_forward<T: Triton, D: Num, const BLOCK_B: i32, const BLOCK_N: i32>(
input_ptr: T::Pointer<D>,
output_ptr: T::Pointer<D>,
B: i32,
N: i32,
stride_ib: i32,
stride_in: i32,
) {
let pid = T::program_id(Axis::X);
let num_pid_n = T::cdiv(N, BLOCK_N);
let pid_b = pid / num_pid_n;
let pid_n = pid % num_pid_n;
let input_desc = T::make_tensor_descriptor(
input_ptr,
&[B, N],
&[stride_ib, stride_in],
&[BLOCK_B, BLOCK_N],
Some(PaddingOption::Zero),
);
let output_desc = T::make_tensor_descriptor(
output_ptr,
&[B, N],
&[N, 1],
&[BLOCK_B, BLOCK_N],
Some(PaddingOption::Zero),
);
let b_off = pid_b * BLOCK_B;
let n_off = pid_n * BLOCK_N;
let tile = T::load_tensor_descriptor(input_desc, &[b_off, n_off]);
T::store_tensor_descriptor(output_desc, &[b_off, n_off], tile);
}
#[kernel]
pub fn flatten_backward<T: Triton, D: Num, const BLOCK_B: i32, const BLOCK_N: i32>(
dy_ptr: T::Pointer<D>,
dx_ptr: T::Pointer<D>,
B: i32,
N: i32,
stride_dxb: i32,
stride_dxn: i32,
) {
let pid = T::program_id(Axis::X);
let num_pid_n = T::cdiv(N, BLOCK_N);
let pid_b = pid / num_pid_n;
let pid_n = pid % num_pid_n;
let dy_desc = T::make_tensor_descriptor(
dy_ptr,
&[B, N],
&[N, 1],
&[BLOCK_B, BLOCK_N],
Some(PaddingOption::Zero),
);
let dx_desc = T::make_tensor_descriptor(
dx_ptr,
&[B, N],
&[stride_dxb, stride_dxn],
&[BLOCK_B, BLOCK_N],
Some(PaddingOption::Zero),
);
let b_off = pid_b * BLOCK_B;
let n_off = pid_n * BLOCK_N;
let tile = T::load_tensor_descriptor(dy_desc, &[b_off, n_off]);
T::store_tensor_descriptor(dx_desc, &[b_off, n_off], tile);
}
impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for FlattenForward<D> {
fn n_activation_inputs(&self) -> usize {
1
}
fn param_shapes(&self, _input_shapes: &[&[usize]], _output_shape: &[usize]) -> Vec<Vec<usize>> {
Vec::new()
}
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 b = output_shape[0] as i32;
let n = output_shape[1] as i32;
visitor.visit_ptr(inputs[0].0);
visitor.visit_ptr(output);
visitor.visit_i32(b);
visitor.visit_i32(n);
visitor.visit_i32(n); visitor.visit_i32(1); }
fn block(&self) -> [u32; 3] {
[128, 1, 1]
}
fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
let pb = output_shape[0].div_ceil(self.block_b as usize);
let pn = output_shape[1].div_ceil(self.block_n as usize);
[(pb * pn) as u32, 1, 1]
}
}
pub struct FlattenOp<'a, T: Num> {
pub forward: FlattenForward<T>,
pub backward: FlattenBackward<T>,
_marker: core::marker::PhantomData<&'a ()>,
}