teeny_kernels/nn/activation/
softmax.rs1#![allow(non_snake_case)]
18
19use teeny_core::dtype::Float;
20use teeny_macros::kernel;
21use teeny_triton::triton::{
22 types::{AddOffsets, Comparison},
23 *,
24};
25
26#[kernel]
39pub fn softmax_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
40 x_ptr: T::Pointer<D>,
41 y_ptr: T::Pointer<D>,
42 _n_rows: i32,
43 n_cols: i32,
44) where
45 T::I32Tensor: types::Tensor<i32, 1>,
46 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
47 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
48{
49 let pid = T::program_id(Axis::X);
50 let row_offset = pid * n_cols;
51 let col_offsets = T::arange(0, BLOCK_SIZE);
52 let offsets = col_offsets + row_offset;
53
54 let x = T::load(
55 x_ptr.add_offsets(offsets),
56 None,
57 None,
58 &[],
59 None,
60 None,
61 None,
62 false,
63 );
64
65 let y = T::softmax(x, None, false, false);
67
68 T::store(y_ptr.add_offsets(offsets), y, None, &[], None, None);
69}
70#[kernel]
88pub fn softmax_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
89 dy_ptr: T::Pointer<D>,
90 y_ptr: T::Pointer<D>,
91 dx_ptr: T::Pointer<D>,
92 _n_rows: i32,
93 n_cols: i32,
94) where
95 T::I32Tensor: types::Tensor<i32, 1>,
96 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
97 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
98{
99 let pid = T::program_id(Axis::X);
100 let row_offset = pid * n_cols;
101 let col_offsets = T::arange(0, BLOCK_SIZE);
102 let offsets = col_offsets + row_offset;
103
104 let dy = T::load(
105 dy_ptr.add_offsets(offsets),
106 None,
107 None,
108 &[],
109 None,
110 None,
111 None,
112 false,
113 );
114 let y = T::load(
115 y_ptr.add_offsets(offsets),
116 None,
117 None,
118 &[],
119 None,
120 None,
121 None,
122 false,
123 );
124
125 let dot = T::sum(y * dy, Some(0), false);
127
128 let dx = y * (dy - dot);
130
131 T::store(dx_ptr.add_offsets(offsets), dx, None, &[], None, None);
132}
133
134impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for SoftmaxForward<D> {
135 fn n_activation_inputs(&self) -> usize {
136 1
137 }
138
139 fn param_shapes(&self, _input_shapes: &[&[usize]], _output_shape: &[usize]) -> Vec<Vec<usize>> {
140 Vec::new()
141 }
142
143 fn pack_args(
144 &self,
145 inputs: &[(teeny_core::model::RawPtr, &[usize])],
146 _params: &[teeny_core::model::RawPtr],
147 output: teeny_core::model::RawPtr,
148 output_shape: &[usize],
149 _output_row_stride: i32,
150 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
151 ) {
152 let n_rows = output_shape[0] as i32;
154 let n_cols = output_shape[1] as i32;
155 visitor.visit_ptr(inputs[0].0);
156 visitor.visit_ptr(output);
157 visitor.visit_i32(n_rows);
158 visitor.visit_i32(n_cols);
159 }
160
161 fn block(&self) -> [u32; 3] {
162 [128, 1, 1]
163 }
164
165 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
167 [output_shape[0] as u32, 1, 1]
168 }
169}
170
171pub struct SoftmaxOp<'a, T: Float> {
172 pub forward: SoftmaxForward<T>,
173 pub backward: SoftmaxBackward<T>,
174 _marker: core::marker::PhantomData<&'a ()>,
175}