1#![allow(non_snake_case)]
29
30use teeny_core::dtype::Float;
31use teeny_macros::kernel;
32use teeny_triton::triton::{
33 types::{AddOffsets, Comparison},
34 *,
35};
36
37#[kernel]
43pub fn layer_norm_forward_inference<T: Triton, D: Float, const BLOCK_N: i32>(
44 x_ptr: T::Pointer<D>,
45 y_ptr: T::Pointer<D>,
46 weight_ptr: T::Pointer<D>,
47 bias_ptr: T::Pointer<D>,
48 _M: i32,
49 N: i32,
50 eps: f32,
51) where
52 T::I32Tensor: types::Tensor<i32, 1>,
53 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
54 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
55{
56 let row = T::program_id(Axis::X);
57 let row_start = row * N;
58
59 let zeros = T::zeros::<D>(&[BLOCK_N]);
61 let zero_1 = T::zeros::<D>(&[1]);
62 let mut sum = zero_1;
63 let mut n_start: i32 = 0;
64 while n_start < N {
65 let col_offs = T::arange(0, BLOCK_N) + n_start;
66 let mask = col_offs.lt(N);
67 let x_tile = T::load(
68 x_ptr.add_offsets(col_offs + row_start),
69 Some(mask),
70 Some(zeros),
71 &[],
72 None,
73 None,
74 None,
75 false,
76 );
77 sum = sum + T::sum(x_tile, None, true);
78 n_start += BLOCK_N;
79 }
80 let n_inv = T::cast::<f32, D>(T::full::<f32>(&[1], 1.0f32 / (N as f32)), None, false);
81 let mean_1 = sum * n_inv;
82 let mean = T::broadcast_to(mean_1, &[BLOCK_N]);
83
84 let mut var_sum = zero_1;
86 n_start = 0;
87 while n_start < N {
88 let col_offs = T::arange(0, BLOCK_N) + n_start;
89 let mask = col_offs.lt(N);
90 let x_tile = T::load(
91 x_ptr.add_offsets(col_offs + row_start),
92 Some(mask),
93 Some(zeros),
94 &[],
95 None,
96 None,
97 None,
98 false,
99 );
100 let diff = T::where_::<D>(mask, x_tile - mean, zeros);
102 var_sum = var_sum + T::sum(diff * diff, None, true);
103 n_start += BLOCK_N;
104 }
105 let eps_t = T::cast::<f32, D>(T::full::<f32>(&[1], eps), None, false);
106 let rstd = T::broadcast_to(T::rsqrt(var_sum * n_inv + eps_t), &[BLOCK_N]);
107
108 n_start = 0;
110 while n_start < N {
111 let col_offs = T::arange(0, BLOCK_N) + n_start;
112 let mask = col_offs.lt(N);
113 let x_tile = T::load(
114 x_ptr.add_offsets(col_offs + row_start),
115 Some(mask),
116 Some(zeros),
117 &[],
118 None,
119 None,
120 None,
121 false,
122 );
123 let gamma = T::load(
124 weight_ptr.add_offsets(col_offs),
125 Some(mask),
126 Some(zeros),
127 &[],
128 None,
129 None,
130 None,
131 false,
132 );
133 let beta = T::load(
134 bias_ptr.add_offsets(col_offs),
135 Some(mask),
136 Some(zeros),
137 &[],
138 None,
139 None,
140 None,
141 false,
142 );
143 let y_tile = (x_tile - mean) * rstd * gamma + beta;
144 T::store(
145 y_ptr.add_offsets(col_offs + row_start),
146 y_tile,
147 Some(mask),
148 &[],
149 None,
150 None,
151 );
152 n_start += BLOCK_N;
153 }
154}
155
156#[cfg(feature = "training")]
162#[kernel]
163pub fn layer_norm_forward<T: Triton, D: Float, const BLOCK_N: i32>(
164 x_ptr: T::Pointer<D>,
165 y_ptr: T::Pointer<D>,
166 weight_ptr: T::Pointer<D>,
167 bias_ptr: T::Pointer<D>,
168 mean_ptr: T::Pointer<D>,
169 rstd_ptr: T::Pointer<D>,
170 _M: i32,
171 N: i32,
172 eps: f32,
173) where
174 T::I32Tensor: types::Tensor<i32, 1>,
175 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
176 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
177{
178 let row = T::program_id(Axis::X);
179 let row_start = row * N;
180 let row_idx = T::arange(0, 1) + row;
181
182 let zeros = T::zeros::<D>(&[BLOCK_N]);
183 let zero_1 = T::zeros::<D>(&[1]);
184 let n_inv = T::cast::<f32, D>(T::full::<f32>(&[1], 1.0f32 / (N as f32)), None, false);
185
186 let mut sum = zero_1;
188 let mut n_start: i32 = 0;
189 while n_start < N {
190 let col_offs = T::arange(0, BLOCK_N) + n_start;
191 let mask = col_offs.lt(N);
192 let x_tile = T::load(
193 x_ptr.add_offsets(col_offs + row_start),
194 Some(mask),
195 Some(zeros),
196 &[],
197 None,
198 None,
199 None,
200 false,
201 );
202 sum = sum + T::sum(x_tile, None, true);
203 n_start += BLOCK_N;
204 }
205 let mean_1 = sum * n_inv;
206 let mean = T::broadcast_to(mean_1, &[BLOCK_N]);
207
208 let mut var_sum = zero_1;
210 n_start = 0;
211 while n_start < N {
212 let col_offs = T::arange(0, BLOCK_N) + n_start;
213 let mask = col_offs.lt(N);
214 let x_tile = T::load(
215 x_ptr.add_offsets(col_offs + row_start),
216 Some(mask),
217 Some(zeros),
218 &[],
219 None,
220 None,
221 None,
222 false,
223 );
224 let diff = T::where_::<D>(mask, x_tile - mean, zeros);
226 var_sum = var_sum + T::sum(diff * diff, None, true);
227 n_start += BLOCK_N;
228 }
229 let eps_t = T::cast::<f32, D>(T::full::<f32>(&[1], eps), None, false);
230 let rstd_1 = T::rsqrt(var_sum * n_inv + eps_t);
231 let rstd = T::broadcast_to(rstd_1, &[BLOCK_N]);
232
233 T::store(mean_ptr.add_offsets(row_idx), mean_1, None, &[], None, None);
234 T::store(rstd_ptr.add_offsets(row_idx), rstd_1, None, &[], None, None);
235
236 n_start = 0;
238 while n_start < N {
239 let col_offs = T::arange(0, BLOCK_N) + n_start;
240 let mask = col_offs.lt(N);
241 let x_tile = T::load(
242 x_ptr.add_offsets(col_offs + row_start),
243 Some(mask),
244 Some(zeros),
245 &[],
246 None,
247 None,
248 None,
249 false,
250 );
251 let gamma = T::load(
252 weight_ptr.add_offsets(col_offs),
253 Some(mask),
254 Some(zeros),
255 &[],
256 None,
257 None,
258 None,
259 false,
260 );
261 let beta = T::load(
262 bias_ptr.add_offsets(col_offs),
263 Some(mask),
264 Some(zeros),
265 &[],
266 None,
267 None,
268 None,
269 false,
270 );
271 let y_tile = (x_tile - mean) * rstd * gamma + beta;
272 T::store(
273 y_ptr.add_offsets(col_offs + row_start),
274 y_tile,
275 Some(mask),
276 &[],
277 None,
278 None,
279 );
280 n_start += BLOCK_N;
281 }
282}
283
284#[cfg(feature = "training")]
300#[kernel]
301pub fn layer_norm_backward<T: Triton, D: Float, const BLOCK_N: i32>(
302 dy_ptr: T::Pointer<D>,
303 x_ptr: T::Pointer<D>,
304 dx_ptr: T::Pointer<D>,
305 weight_ptr: T::Pointer<D>,
306 dweight_ptr: T::Pointer<D>,
307 dbias_ptr: T::Pointer<D>,
308 mean_ptr: T::Pointer<D>,
309 rstd_ptr: T::Pointer<D>,
310 _M: i32,
311 N: i32,
312) where
313 T::I32Tensor: types::Tensor<i32, 1>,
314 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
315 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
316{
317 let row = T::program_id(Axis::X);
318 let row_start = row * N;
319 let row_idx = T::arange(0, 1) + row;
320
321 let zeros = T::zeros::<D>(&[BLOCK_N]);
322 let zero_1 = T::zeros::<D>(&[1]);
323 let n_inv = T::cast::<f32, D>(T::full::<f32>(&[1], 1.0f32 / (N as f32)), None, false);
324
325 let rstd_1 = T::load(
326 rstd_ptr.add_offsets(row_idx),
327 None,
328 None,
329 &[],
330 None,
331 None,
332 None,
333 false,
334 );
335 let mean_1 = T::load(
336 mean_ptr.add_offsets(row_idx),
337 None,
338 None,
339 &[],
340 None,
341 None,
342 None,
343 false,
344 );
345 let rstd = T::broadcast_to(rstd_1, &[BLOCK_N]);
346 let mean = T::broadcast_to(mean_1, &[BLOCK_N]);
347
348 let mut sum_dy_gamma = zero_1;
350 let mut sum_dy_gamma_xhat = zero_1;
351 let mut n_start: i32 = 0;
352 while n_start < N {
353 let col_offs = T::arange(0, BLOCK_N) + n_start;
354 let mask = col_offs.lt(N);
355 let x_tile = T::load(
356 x_ptr.add_offsets(col_offs + row_start),
357 Some(mask),
358 Some(zeros),
359 &[],
360 None,
361 None,
362 None,
363 false,
364 );
365 let dy_tile = T::load(
366 dy_ptr.add_offsets(col_offs + row_start),
367 Some(mask),
368 Some(zeros),
369 &[],
370 None,
371 None,
372 None,
373 false,
374 );
375 let gamma = T::load(
376 weight_ptr.add_offsets(col_offs),
377 Some(mask),
378 Some(zeros),
379 &[],
380 None,
381 None,
382 None,
383 false,
384 );
385 let xhat = (x_tile - mean) * rstd;
386 sum_dy_gamma = sum_dy_gamma + T::sum(dy_tile * gamma, None, true);
387 sum_dy_gamma_xhat = sum_dy_gamma_xhat + T::sum(dy_tile * gamma * xhat, None, true);
388 n_start += BLOCK_N;
389 }
390 let c1 = T::broadcast_to(sum_dy_gamma * n_inv, &[BLOCK_N]);
391 let c2 = T::broadcast_to(sum_dy_gamma_xhat * n_inv, &[BLOCK_N]);
392
393 n_start = 0;
395 while n_start < N {
396 let col_offs = T::arange(0, BLOCK_N) + n_start;
397 let mask = col_offs.lt(N);
398 let x_tile = T::load(
399 x_ptr.add_offsets(col_offs + row_start),
400 Some(mask),
401 Some(zeros),
402 &[],
403 None,
404 None,
405 None,
406 false,
407 );
408 let dy_tile = T::load(
409 dy_ptr.add_offsets(col_offs + row_start),
410 Some(mask),
411 Some(zeros),
412 &[],
413 None,
414 None,
415 None,
416 false,
417 );
418 let gamma = T::load(
419 weight_ptr.add_offsets(col_offs),
420 Some(mask),
421 Some(zeros),
422 &[],
423 None,
424 None,
425 None,
426 false,
427 );
428 let dw_old = T::load(
429 dweight_ptr.add_offsets(col_offs),
430 Some(mask),
431 Some(zeros),
432 &[],
433 None,
434 None,
435 None,
436 false,
437 );
438 let db_old = T::load(
439 dbias_ptr.add_offsets(col_offs),
440 Some(mask),
441 Some(zeros),
442 &[],
443 None,
444 None,
445 None,
446 false,
447 );
448
449 let xhat = (x_tile - mean) * rstd;
450 let dx_tile = rstd * gamma * (dy_tile - c1 - xhat * c2);
451
452 T::store(
453 dx_ptr.add_offsets(col_offs + row_start),
454 dx_tile,
455 Some(mask),
456 &[],
457 None,
458 None,
459 );
460 T::store(
461 dweight_ptr.add_offsets(col_offs),
462 dw_old + dy_tile * xhat,
463 Some(mask),
464 &[],
465 None,
466 None,
467 );
468 T::store(
469 dbias_ptr.add_offsets(col_offs),
470 db_old + dy_tile,
471 Some(mask),
472 &[],
473 None,
474 None,
475 );
476 n_start += BLOCK_N;
477 }
478}
479
480pub struct LayerNormForwardInferenceRuntimeOp<D: Float + Send + Sync + 'static> {
487 fwd: LayerNormForwardInference<D>,
488 #[allow(dead_code)]
489 block_n: i32,
490 eps: f32,
491}
492
493impl<D: Float + Send + Sync + 'static> LayerNormForwardInferenceRuntimeOp<D> {
494 pub fn new(block_n: i32, eps: f32) -> Self {
495 Self {
496 fwd: LayerNormForwardInference::<D>::new(block_n),
497 block_n,
498 eps,
499 }
500 }
501
502 pub fn forward_source(&self) -> &str {
503 &self.fwd.source
504 }
505 pub fn kernel_name(&self) -> &str {
506 self.fwd.name
507 }
508}
509
510impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp
511 for LayerNormForwardInferenceRuntimeOp<D>
512{
513 fn n_activation_inputs(&self) -> usize {
514 1
515 }
516
517 fn param_shapes(&self, input_shapes: &[&[usize]], _output_shape: &[usize]) -> Vec<Vec<usize>> {
518 let n = *input_shapes[0].last().unwrap();
520 vec![vec![n], vec![n]]
521 }
522
523 fn param_names(&self) -> &'static [&'static str] {
524 &["weight", "bias"]
525 }
526
527 fn pack_args(
528 &self,
529 inputs: &[(teeny_core::model::RawPtr, &[usize])],
530 params: &[teeny_core::model::RawPtr],
531 output: teeny_core::model::RawPtr,
532 _output_shape: &[usize],
533 _output_row_stride: i32,
534 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
535 ) {
536 let shape = inputs[0].1;
537 let n = *shape.last().unwrap() as i32;
538 let total: usize = shape.iter().product();
539 let m = (total as i32) / n;
540
541 visitor.visit_ptr(inputs[0].0); visitor.visit_ptr(output); visitor.visit_ptr(params[0]); visitor.visit_ptr(params[1]); visitor.visit_i32(m); visitor.visit_i32(n); visitor.visit_f32(self.eps); }
549
550 fn block(&self) -> [u32; 3] {
551 [1, 1, 1]
552 }
553
554 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
555 let n = *output_shape.last().unwrap();
557 let total: usize = output_shape.iter().product();
558 let m = total / n;
559 [m as u32, 1, 1]
560 }
561
562 fn has_backward(&self) -> bool {
563 false
564 }
565}