1use teeny_core::dtype::{AddOffsets, Comparison, Num, Tensor};
18use teeny_macros::kernel;
19use teeny_triton::triton::{Axis, InputPrecision, PaddingOption, Triton};
20
21#[kernel]
22pub fn linear_forward<
23 T: Triton,
24 D: Num,
25 const USE_BIAS: bool,
26 const BLOCK_M: i32,
27 const BLOCK_N: i32,
28 const BLOCK_K: i32,
29 const GROUP_M: i32,
30>(
31 x_ptr: T::Pointer<D>,
32 w_ptr: T::Pointer<D>,
33 b_ptr: T::Pointer<D>,
34 y_ptr: T::Pointer<D>,
35 M: i32,
36 N: i32,
37 K: i32,
38 stride_xm: i32,
39 stride_xk: i32,
40 stride_wn: i32,
41 stride_wk: i32,
42 stride_ym: i32,
43 stride_yn: i32,
44) where
45 T::I32Tensor: 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 num_pid_m = T::cdiv(M, BLOCK_M);
51 let num_pid_n = T::cdiv(N, BLOCK_N);
52 let num_pid_in_group = GROUP_M * num_pid_n;
53 let group_id = pid / num_pid_in_group;
54 let first_pid_m = group_id * GROUP_M;
55 let remaining_m = num_pid_m - first_pid_m;
56 let group_size_m = if remaining_m < GROUP_M {
57 remaining_m
58 } else {
59 GROUP_M
60 };
61 let pid_in_group = pid % num_pid_in_group;
62 let pid_m = first_pid_m + (pid_in_group % group_size_m);
63 let pid_n = pid_in_group / group_size_m;
64
65 let x_desc = T::make_tensor_descriptor(
66 x_ptr,
67 &[M, K],
68 &[stride_xm, stride_xk],
69 &[BLOCK_M, BLOCK_K],
70 Some(PaddingOption::Zero),
71 );
72 let w_desc = T::make_tensor_descriptor(
73 w_ptr,
74 &[N, K],
75 &[stride_wn, stride_wk],
76 &[BLOCK_N, BLOCK_K],
77 Some(PaddingOption::Zero),
78 );
79
80 let mut acc = T::zeros::<D>(&[BLOCK_M, BLOCK_N]);
81 let k_tiles = T::cdiv(K, BLOCK_K);
82 for k in 0..k_tiles {
83 let x = T::load_tensor_descriptor(x_desc, &[pid_m * BLOCK_M, k * BLOCK_K]);
84 let w = T::load_tensor_descriptor(w_desc, &[pid_n * BLOCK_N, k * BLOCK_K]);
85 let w_t = T::trans(w, &[1, 0]);
86 acc = T::dot::<D, D>(x, w_t, Some(acc), InputPrecision::TF32, None);
87 }
88
89 if USE_BIAS {
90 let offs_bn = T::arange(0, BLOCK_N) + pid_n * BLOCK_N;
91 let bias_mask = offs_bn.lt(N);
92 let bias = T::load(
93 b_ptr.add_offsets(offs_bn),
94 Some(bias_mask),
95 Some(T::zeros::<D>(&[BLOCK_N])),
96 &[],
97 None,
98 None,
99 None,
100 false,
101 );
102 let bias = T::expand_dims(bias, 0);
103 let bias = T::broadcast_to(bias, &[BLOCK_M, BLOCK_N]);
104 acc = acc + bias;
105 }
106
107 let y_desc = T::make_tensor_descriptor(
108 y_ptr,
109 &[M, N],
110 &[stride_ym, stride_yn],
111 &[BLOCK_M, BLOCK_N],
112 Some(PaddingOption::Zero),
113 );
114
115 T::store_tensor_descriptor(y_desc, &[pid_m * BLOCK_M, pid_n * BLOCK_N], acc);
116}
117
118#[kernel]
119pub fn linear_backward<
120 T: Triton,
121 D: Num,
122 const USE_BIAS: bool,
123 const BLOCK_M: i32,
124 const BLOCK_N: i32,
125 const BLOCK_K: i32,
126 const GROUP_M: i32,
127>(
128 x_ptr: T::Pointer<D>,
129 w_ptr: T::Pointer<D>,
130 dy_ptr: T::Pointer<D>,
131 dx_ptr: T::Pointer<D>,
132 dw_ptr: T::Pointer<D>,
133 db_ptr: T::Pointer<D>,
134 M: i32,
135 N: i32,
136 K: i32,
137 stride_xm: i32,
138 stride_xk: i32,
139 stride_wk: i32,
140 stride_wn: i32,
141 stride_dym: i32,
142 stride_dyn: i32,
143 stride_dxm: i32,
144 stride_dxk: i32,
145 stride_dwk: i32,
146 stride_dwn: i32,
147 _stride_dbn: i32,
148) where
149 T::I32Tensor: Tensor<i32, 1>,
150 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
151 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
152{
153 let pid = T::program_id(Axis::X);
162 let num_pid_k = T::cdiv(K, BLOCK_K);
163 let num_pid_n = T::cdiv(N, BLOCK_N);
164 let pid_k = pid % num_pid_k;
165 let pid_tmp = pid / num_pid_k;
166 let pid_n = pid_tmp % num_pid_n;
167 let pid_m = pid_tmp / num_pid_n;
168
169 let x_desc = T::make_tensor_descriptor(
170 x_ptr,
171 &[M, K],
172 &[stride_xm, stride_xk],
173 &[BLOCK_M, BLOCK_K],
174 Some(PaddingOption::Zero),
175 );
176 let w_desc = T::make_tensor_descriptor(
177 w_ptr,
178 &[N, K],
179 &[stride_wn, stride_wk],
180 &[BLOCK_N, BLOCK_K],
181 Some(PaddingOption::Zero),
182 );
183 let dy_desc = T::make_tensor_descriptor(
184 dy_ptr,
185 &[M, N],
186 &[stride_dym, stride_dyn],
187 &[BLOCK_M, BLOCK_N],
188 Some(PaddingOption::Zero),
189 );
190
191 let dx_desc = T::make_tensor_descriptor(
196 dx_ptr,
197 &[M, K],
198 &[stride_dxm, stride_dxk],
199 &[BLOCK_M, BLOCK_K],
200 Some(PaddingOption::Zero),
201 );
202 if pid_n == 0 {
203 let n_tiles = T::cdiv(N, BLOCK_N);
204 let mut acc_dx = T::zeros::<D>(&[BLOCK_M, BLOCK_K]);
205 for n in 0..n_tiles {
206 let dy = T::load_tensor_descriptor(dy_desc, &[pid_m * BLOCK_M, n * BLOCK_N]);
207 let w = T::load_tensor_descriptor(w_desc, &[n * BLOCK_N, pid_k * BLOCK_K]);
208 acc_dx = T::dot::<D, D>(dy, w, Some(acc_dx), InputPrecision::TF32, None);
209 }
210 T::store_tensor_descriptor(dx_desc, &[pid_m * BLOCK_M, pid_k * BLOCK_K], acc_dx);
211 }
212
213 let dw_desc = T::make_tensor_descriptor(
218 dw_ptr,
219 &[N, K],
220 &[stride_dwn, stride_dwk],
221 &[BLOCK_N, BLOCK_K],
222 Some(PaddingOption::Zero),
223 );
224 if pid_m == 0 {
225 let m_tiles = T::cdiv(M, BLOCK_M);
226 let mut acc_dw = T::zeros::<D>(&[BLOCK_N, BLOCK_K]);
227 for m in 0..m_tiles {
228 let dy = T::load_tensor_descriptor(dy_desc, &[m * BLOCK_M, pid_n * BLOCK_N]);
229 let x = T::load_tensor_descriptor(x_desc, &[m * BLOCK_M, pid_k * BLOCK_K]);
230 let dy_t = T::trans(dy, &[1, 0]);
231 acc_dw = T::dot::<D, D>(dy_t, x, Some(acc_dw), InputPrecision::TF32, None);
232 }
233 T::store_tensor_descriptor(dw_desc, &[pid_n * BLOCK_N, pid_k * BLOCK_K], acc_dw);
234
235 if USE_BIAS && pid_k == 0 {
237 let offs_bn = T::arange(0, BLOCK_N) + pid_n * BLOCK_N;
238 let bias_mask = offs_bn.lt(N);
239 let mut acc_db = T::zeros::<D>(&[BLOCK_N]);
240 for m in 0..m_tiles {
241 let dy = T::load_tensor_descriptor(dy_desc, &[m * BLOCK_M, pid_n * BLOCK_N]);
242 let sum = T::sum::<D>(dy, Some(0), false);
243 acc_db = acc_db + sum;
244 }
245 let db_ptr_tile = db_ptr.add_offsets(offs_bn);
246 T::store(db_ptr_tile, acc_db, Some(bias_mask), &[], None, None);
247 }
248 }
249}
250
251impl<D: Num + Send + Sync + 'static> teeny_core::model::RuntimeOp for LinearForward<D> {
252 fn n_activation_inputs(&self) -> usize {
253 1
254 }
255
256 fn param_shapes(&self, input_shapes: &[&[usize]], output_shape: &[usize]) -> Vec<Vec<usize>> {
257 let k = input_shapes[0][1];
259 let n = output_shape[1];
260 if self.use_bias {
262 vec![vec![n, k], vec![n]]
263 } else {
264 vec![vec![n, k]]
265 }
266 }
267
268 fn forward_output_row_stride(&self, output_shape: &[usize]) -> usize {
270 let n = output_shape.last().copied().unwrap_or(1);
271 let align = 16 / core::mem::size_of::<D>();
272 n.next_multiple_of(align)
273 }
274
275 fn pack_args(
276 &self,
277 inputs: &[(teeny_core::model::RawPtr, &[usize])],
278 params: &[teeny_core::model::RawPtr],
279 output: teeny_core::model::RawPtr,
280 output_shape: &[usize],
281 output_row_stride: i32,
282 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
283 ) {
284 let m = output_shape[0] as i32;
287 let n = output_shape[1] as i32;
288 let k = inputs[0].1[1] as i32;
289 let b_ptr = if self.use_bias {
290 params[1]
291 } else {
292 core::ptr::null_mut()
293 };
294 visitor.visit_ptr(inputs[0].0); visitor.visit_ptr(params[0]); visitor.visit_ptr(b_ptr); visitor.visit_ptr(output); visitor.visit_i32(m); visitor.visit_i32(n); visitor.visit_i32(k); visitor.visit_i32(k); visitor.visit_i32(1); visitor.visit_i32(k); visitor.visit_i32(1); visitor.visit_i32(output_row_stride); visitor.visit_i32(1); }
308
309 fn block(&self) -> [u32; 3] {
310 [128, 1, 1]
311 }
312
313 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
314 let pm = output_shape[0].div_ceil(self.block_m as usize);
316 let pn = output_shape[1].div_ceil(self.block_n as usize);
317 [(pm * pn) as u32, 1, 1]
318 }
319
320 #[cfg(feature = "training")]
321 fn has_backward(&self) -> bool {
322 true
323 }
324
325 #[cfg(feature = "training")]
328 fn backward_grad_output_row_stride(&self, output_shape: &[usize]) -> usize {
329 let n = output_shape[output_shape.len() - 1];
330 let align = 16 / core::mem::size_of::<D>();
331 n.next_multiple_of(align)
332 }
333
334 #[cfg(feature = "training")]
340 #[allow(clippy::too_many_arguments)]
341 fn pack_backward_args(
342 &self,
343 inputs: &[(teeny_core::model::RawPtr, &[usize])],
344 params: &[teeny_core::model::RawPtr],
345 _output: teeny_core::model::RawPtr,
346 output_shape: &[usize],
347 grad_output: teeny_core::model::RawPtr,
348 grad_output_row_stride: i32,
349 grad_inputs: &[teeny_core::model::RawPtr],
350 grad_params: &[teeny_core::model::RawPtr],
351 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
352 ) {
353 let m = output_shape[0] as i32;
354 let n = output_shape[1] as i32;
355 let k = inputs[0].1[1] as i32;
356 let db_ptr = if self.use_bias {
357 grad_params[1]
358 } else {
359 core::ptr::null_mut()
360 };
361
362 visitor.visit_ptr(inputs[0].0); visitor.visit_ptr(params[0]); visitor.visit_ptr(grad_output); visitor.visit_ptr(grad_inputs[0]); visitor.visit_ptr(grad_params[0]); visitor.visit_ptr(db_ptr); visitor.visit_i32(m); visitor.visit_i32(n); visitor.visit_i32(k); visitor.visit_i32(k); visitor.visit_i32(1); visitor.visit_i32(1); visitor.visit_i32(k); visitor.visit_i32(grad_output_row_stride); visitor.visit_i32(1); visitor.visit_i32(k); visitor.visit_i32(1); visitor.visit_i32(1); visitor.visit_i32(k); visitor.visit_i32(1); }
383
384 #[cfg(feature = "training")]
385 fn backward_block(&self) -> [u32; 3] {
386 [128, 1, 1]
387 }
388
389 #[cfg(feature = "training")]
391 fn backward_grid(&self, input_shapes: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
392 let m = output_shape[0].div_ceil(self.block_m as usize);
393 let n = output_shape[1].div_ceil(self.block_n as usize);
394 let k = input_shapes[0][1].div_ceil(self.block_k as usize);
395 [(m * n * k) as u32, 1, 1]
396 }
397}
398
399pub struct LinearOp<'a, T: Num> {
400 pub forward: LinearForward<T>,
401 pub backward: LinearBackward<T>,
402 _marker: core::marker::PhantomData<&'a ()>,
403}