1#![allow(non_snake_case)]
34
35use teeny_core::dtype::Float;
36use teeny_macros::kernel;
37use teeny_triton::triton::{
38 types::{AddOffsets, Comparison},
39 *,
40};
41
42#[kernel]
48pub fn batch_norm_forward_inference<T: Triton, D: Float, const BLOCK_N: i32>(
49 x_ptr: T::Pointer<D>,
50 y_ptr: T::Pointer<D>,
51 weight_ptr: T::Pointer<D>,
52 bias_ptr: T::Pointer<D>,
53 running_mean_ptr: T::Pointer<D>,
54 running_var_ptr: T::Pointer<D>,
55 N: i32,
56 C: i32,
57 eps: f32,
58) where
59 T::I32Tensor: types::Tensor<i32, 1>,
60 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
61 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
62{
63 let c = T::program_id(Axis::X);
64 let c_idx = T::arange(0, 1) + c;
65
66 let mean = T::broadcast_to(
68 T::load(
69 running_mean_ptr.add_offsets(c_idx),
70 None,
71 None,
72 &[],
73 None,
74 None,
75 None,
76 false,
77 ),
78 &[BLOCK_N],
79 );
80 let var = T::load(
81 running_var_ptr.add_offsets(c_idx),
82 None,
83 None,
84 &[],
85 None,
86 None,
87 None,
88 false,
89 );
90 let rstd = T::broadcast_to(
91 T::rsqrt(var + T::cast::<f32, D>(T::full::<f32>(&[1], eps), None, false)),
92 &[BLOCK_N],
93 );
94 let gamma = T::broadcast_to(
95 T::load(
96 weight_ptr.add_offsets(c_idx),
97 None,
98 None,
99 &[],
100 None,
101 None,
102 None,
103 false,
104 ),
105 &[BLOCK_N],
106 );
107 let beta = T::broadcast_to(
108 T::load(
109 bias_ptr.add_offsets(c_idx),
110 None,
111 None,
112 &[],
113 None,
114 None,
115 None,
116 false,
117 ),
118 &[BLOCK_N],
119 );
120
121 let zeros = T::zeros::<D>(&[BLOCK_N]);
123 let mut n_start: i32 = 0;
124 while n_start < N {
125 let offsets_n = T::arange(0, BLOCK_N) + n_start;
126 let mask = offsets_n.lt(N);
127 let elem_offsets = offsets_n * C + c;
128
129 let x_tile = T::load(
130 x_ptr.add_offsets(elem_offsets),
131 Some(mask),
132 Some(zeros),
133 &[],
134 None,
135 None,
136 None,
137 false,
138 );
139 let y_tile = gamma * (x_tile - mean) * rstd + beta;
140
141 T::store(
142 y_ptr.add_offsets(elem_offsets),
143 y_tile,
144 Some(mask),
145 &[],
146 None,
147 None,
148 );
149
150 n_start += BLOCK_N;
151 }
152}
153
154#[cfg(feature = "training")]
164#[kernel]
165pub fn batch_norm_stats_forward<T: Triton, D: Float, const BLOCK_N: i32>(
166 x_ptr: T::Pointer<D>,
167 mean_ptr: T::Pointer<D>,
168 rstd_ptr: T::Pointer<D>,
169 running_mean_ptr: T::Pointer<D>,
170 running_var_ptr: T::Pointer<D>,
171 N: i32,
172 C: i32,
173 eps: f32,
174 momentum: f32,
175) where
176 T::I32Tensor: types::Tensor<i32, 1>,
177 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
178 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
179{
180 let c = T::program_id(Axis::X);
181
182 let zeros_blk = T::zeros::<D>(&[BLOCK_N]);
186 let mut acc_sum = zeros_blk;
187 let mut acc_sum_sq = zeros_blk;
188 let mut n_start: i32 = 0;
189
190 while n_start < N {
191 let offsets_n = T::arange(0, BLOCK_N) + n_start;
192 let mask = offsets_n.lt(N);
193 let elem_offsets = offsets_n * C + c;
194
195 let x_tile = T::load(
196 x_ptr.add_offsets(elem_offsets),
197 Some(mask),
198 Some(zeros_blk),
199 &[],
200 None,
201 None,
202 None,
203 false,
204 );
205 acc_sum = acc_sum + x_tile;
206 acc_sum_sq = acc_sum_sq + x_tile * x_tile;
207
208 n_start += BLOCK_N;
209 }
210
211 let sum = T::sum(acc_sum, None, true);
213 let sum_sq = T::sum(acc_sum_sq, None, true);
214
215 let n_inv = T::cast::<f32, D>(T::full::<f32>(&[1], 1.0f32 / (N as f32)), None, false);
217 let mean_1 = sum * n_inv;
218 let var_1 = sum_sq * n_inv - mean_1 * mean_1;
219 let rstd_1 = T::rsqrt(var_1 + T::cast::<f32, D>(T::full::<f32>(&[1], eps), None, false));
220
221 let c_idx = T::arange(0, 1) + c;
223 T::store(mean_ptr.add_offsets(c_idx), mean_1, None, &[], None, None);
224 T::store(rstd_ptr.add_offsets(c_idx), rstd_1, None, &[], None, None);
225
226 let m = T::cast::<f32, D>(T::full::<f32>(&[1], momentum), None, false);
228 let one_m = T::cast::<f32, D>(T::full::<f32>(&[1], 1.0f32 - momentum), None, false);
229 let running_mean_old = T::load(
230 running_mean_ptr.add_offsets(c_idx),
231 None,
232 None,
233 &[],
234 None,
235 None,
236 None,
237 false,
238 );
239 let running_var_old = T::load(
240 running_var_ptr.add_offsets(c_idx),
241 None,
242 None,
243 &[],
244 None,
245 None,
246 None,
247 false,
248 );
249
250 T::store(
251 running_mean_ptr.add_offsets(c_idx),
252 one_m * running_mean_old + m * mean_1,
253 None,
254 &[],
255 None,
256 None,
257 );
258 T::store(
259 running_var_ptr.add_offsets(c_idx),
260 one_m * running_var_old + m * var_1,
261 None,
262 &[],
263 None,
264 None,
265 );
266}
267
268#[cfg(feature = "training")]
274#[kernel]
275pub fn batch_norm_normalize_forward<T: Triton, D: Float, const BLOCK_N: i32>(
276 x_ptr: T::Pointer<D>,
277 y_ptr: T::Pointer<D>,
278 weight_ptr: T::Pointer<D>,
279 bias_ptr: T::Pointer<D>,
280 mean_ptr: T::Pointer<D>,
281 rstd_ptr: T::Pointer<D>,
282 N: i32,
283 C: i32,
284) where
285 T::I32Tensor: types::Tensor<i32, 1>,
286 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
287 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
288{
289 let c = T::program_id(Axis::X);
290 let c_idx = T::arange(0, 1) + c;
291
292 let mean = T::broadcast_to(
294 T::load(
295 mean_ptr.add_offsets(c_idx),
296 None,
297 None,
298 &[],
299 None,
300 None,
301 None,
302 false,
303 ),
304 &[BLOCK_N],
305 );
306 let rstd = T::broadcast_to(
307 T::load(
308 rstd_ptr.add_offsets(c_idx),
309 None,
310 None,
311 &[],
312 None,
313 None,
314 None,
315 false,
316 ),
317 &[BLOCK_N],
318 );
319 let gamma = T::broadcast_to(
320 T::load(
321 weight_ptr.add_offsets(c_idx),
322 None,
323 None,
324 &[],
325 None,
326 None,
327 None,
328 false,
329 ),
330 &[BLOCK_N],
331 );
332 let beta = T::broadcast_to(
333 T::load(
334 bias_ptr.add_offsets(c_idx),
335 None,
336 None,
337 &[],
338 None,
339 None,
340 None,
341 false,
342 ),
343 &[BLOCK_N],
344 );
345
346 let zeros = T::zeros::<D>(&[BLOCK_N]);
347 let mut n_start: i32 = 0;
348 while n_start < N {
349 let offsets_n = T::arange(0, BLOCK_N) + n_start;
350 let mask = offsets_n.lt(N);
351 let elem_offsets = offsets_n * C + c;
352
353 let x_tile = T::load(
354 x_ptr.add_offsets(elem_offsets),
355 Some(mask),
356 Some(zeros),
357 &[],
358 None,
359 None,
360 None,
361 false,
362 );
363 let y_tile = gamma * (x_tile - mean) * rstd + beta;
364
365 T::store(
366 y_ptr.add_offsets(elem_offsets),
367 y_tile,
368 Some(mask),
369 &[],
370 None,
371 None,
372 );
373
374 n_start += BLOCK_N;
375 }
376}
377
378#[cfg(feature = "training")]
385pub struct BatchNormStatsRuntimeOp<D: teeny_core::dtype::Float + Send + Sync + 'static> {
386 pub block_n: i32,
387 pub eps: f32,
388 pub momentum: f32,
389 _phantom: core::marker::PhantomData<D>,
390}
391
392#[cfg(feature = "training")]
393impl<D: teeny_core::dtype::Float + Send + Sync + 'static> BatchNormStatsRuntimeOp<D> {
394 pub fn new(block_n: i32, eps: f32, momentum: f32) -> Self {
395 Self {
396 block_n,
397 eps,
398 momentum,
399 _phantom: core::marker::PhantomData,
400 }
401 }
402}
403
404#[cfg(feature = "training")]
405impl<D: teeny_core::dtype::Float + Send + Sync + 'static> teeny_core::model::RuntimeOp
406 for BatchNormStatsRuntimeOp<D>
407{
408 fn n_activation_inputs(&self) -> usize {
409 1
410 }
411
412 fn param_shapes(&self, input_shapes: &[&[usize]], _output_shape: &[usize]) -> Vec<Vec<usize>> {
413 let c = input_shapes[0][1];
414 vec![vec![c], vec![c]]
415 }
416
417 fn pack_args(
418 &self,
419 inputs: &[(teeny_core::model::RawPtr, &[usize])],
420 params: &[teeny_core::model::RawPtr],
421 output: teeny_core::model::RawPtr,
422 output_shape: &[usize],
423 _output_row_stride: i32,
424 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
425 ) {
426 let c = output_shape[0] / 2;
427 let n_total: usize = inputs[0].1.iter().product();
428 let n = (n_total / c) as i32;
429 let mean_ptr = output;
430 let rstd_ptr = unsafe { (output as *mut D).add(c) } as teeny_core::model::RawPtr;
431 visitor.visit_ptr(inputs[0].0);
432 visitor.visit_ptr(mean_ptr);
433 visitor.visit_ptr(rstd_ptr);
434 visitor.visit_ptr(params[0]);
435 visitor.visit_ptr(params[1]);
436 visitor.visit_i32(n);
437 visitor.visit_i32(c as i32);
438 visitor.visit_f32(self.eps);
439 visitor.visit_f32(self.momentum);
440 }
441
442 fn block(&self) -> [u32; 3] {
443 [1, 1, 1]
444 }
445
446 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
447 let c = output_shape[0] / 2;
448 [c as u32, 1, 1]
449 }
450}
451
452#[cfg(feature = "training")]
457pub struct BatchNormNormalizeRuntimeOp<D: teeny_core::dtype::Float + Send + Sync + 'static> {
458 pub block_n: i32,
459 bwd_source: String,
460 _phantom: core::marker::PhantomData<D>,
461}
462
463#[cfg(feature = "training")]
464impl<D: teeny_core::dtype::Float + Send + Sync + 'static> BatchNormNormalizeRuntimeOp<D> {
465 pub fn new(block_n: i32) -> Self {
466 Self {
467 block_n,
468 bwd_source: BatchNormBackward::<D>::new(block_n).source,
469 _phantom: core::marker::PhantomData,
470 }
471 }
472
473 pub fn backward_source(&self) -> &str {
474 &self.bwd_source
475 }
476}
477
478#[cfg(feature = "training")]
479impl<D: teeny_core::dtype::Float + Send + Sync + 'static> teeny_core::model::RuntimeOp
480 for BatchNormNormalizeRuntimeOp<D>
481{
482 fn n_activation_inputs(&self) -> usize {
483 2
484 }
485
486 fn param_shapes(&self, input_shapes: &[&[usize]], _output_shape: &[usize]) -> Vec<Vec<usize>> {
487 let c = input_shapes[1][0] / 2;
488 vec![vec![c], vec![c]]
489 }
490
491 fn param_names(&self) -> &'static [&'static str] {
492 &["weight", "bias"]
493 }
494
495 fn pack_args(
496 &self,
497 inputs: &[(teeny_core::model::RawPtr, &[usize])],
498 params: &[teeny_core::model::RawPtr],
499 output: teeny_core::model::RawPtr,
500 _output_shape: &[usize],
501 _output_row_stride: i32,
502 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
503 ) {
504 let c = inputs[1].1[0] / 2;
505 let n_total: usize = inputs[0].1.iter().product();
506 let n = (n_total / c) as i32;
507 let mean_ptr = inputs[1].0;
508 let rstd_ptr = unsafe { (inputs[1].0 as *mut D).add(c) } as teeny_core::model::RawPtr;
509 visitor.visit_ptr(inputs[0].0);
510 visitor.visit_ptr(output);
511 visitor.visit_ptr(params[0]);
512 visitor.visit_ptr(params[1]);
513 visitor.visit_ptr(mean_ptr);
514 visitor.visit_ptr(rstd_ptr);
515 visitor.visit_i32(n);
516 visitor.visit_i32(c as i32);
517 }
518
519 fn block(&self) -> [u32; 3] {
520 [1, 1, 1]
521 }
522
523 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
524 let c = output_shape.get(1).copied().unwrap_or(output_shape[0]);
525 [c as u32, 1, 1]
526 }
527
528 fn has_backward(&self) -> bool {
529 true
530 }
531
532 fn pack_backward_args(
533 &self,
534 inputs: &[(teeny_core::model::RawPtr, &[usize])],
535 params: &[teeny_core::model::RawPtr],
536 _output: teeny_core::model::RawPtr,
537 _output_shape: &[usize],
538 grad_output: teeny_core::model::RawPtr,
539 _grad_output_row_stride: i32,
540 grad_inputs: &[teeny_core::model::RawPtr],
541 grad_params: &[teeny_core::model::RawPtr],
542 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
543 ) {
544 let c = inputs[1].1[0] / 2;
549 let n_total: usize = inputs[0].1.iter().product();
550 let n = (n_total / c) as i32;
551 let mean_ptr = inputs[1].0;
552 let rstd_ptr = unsafe { (inputs[1].0 as *mut D).add(c) } as teeny_core::model::RawPtr;
553 visitor.visit_ptr(grad_output); visitor.visit_ptr(inputs[0].0); visitor.visit_ptr(grad_inputs[0]); visitor.visit_ptr(params[0]); visitor.visit_ptr(mean_ptr); visitor.visit_ptr(rstd_ptr); visitor.visit_ptr(grad_params[0]); visitor.visit_ptr(grad_params[1]); visitor.visit_i32(n);
562 visitor.visit_i32(c as i32);
563 }
564
565 fn backward_block(&self) -> [u32; 3] {
566 [1, 1, 1]
567 }
568
569 fn backward_grid(&self, input_shapes: &[&[usize]], _output_shape: &[usize]) -> [u32; 3] {
570 let c = input_shapes[1][0] / 2;
572 [c as u32, 1, 1]
573 }
574}
575
576#[kernel]
586pub fn batch_norm_2d_nchw_forward_inference<T: Triton, D: Float, const BLOCK_HW: i32>(
587 x_ptr: T::Pointer<D>,
588 y_ptr: T::Pointer<D>,
589 weight_ptr: T::Pointer<D>,
590 bias_ptr: T::Pointer<D>,
591 running_mean_ptr: T::Pointer<D>,
592 running_var_ptr: T::Pointer<D>,
593 C: i32,
594 HW: i32,
595 eps: f32,
596) where
597 T::I32Tensor: types::Tensor<i32, 1>,
598 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
599 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
600{
601 let c = T::program_id(Axis::X);
602 let b = T::program_id(Axis::Y);
603 let c_idx = T::arange(0, 1) + c;
604
605 let mean = T::broadcast_to(
607 T::load(
608 running_mean_ptr.add_offsets(c_idx),
609 None,
610 None,
611 &[],
612 None,
613 None,
614 None,
615 false,
616 ),
617 &[BLOCK_HW],
618 );
619 let var = T::load(
620 running_var_ptr.add_offsets(c_idx),
621 None,
622 None,
623 &[],
624 None,
625 None,
626 None,
627 false,
628 );
629 let rstd = T::broadcast_to(
630 T::rsqrt(var + T::cast::<f32, D>(T::full::<f32>(&[1], eps), None, false)),
631 &[BLOCK_HW],
632 );
633 let gamma = T::broadcast_to(
634 T::load(
635 weight_ptr.add_offsets(c_idx),
636 None,
637 None,
638 &[],
639 None,
640 None,
641 None,
642 false,
643 ),
644 &[BLOCK_HW],
645 );
646 let beta = T::broadcast_to(
647 T::load(
648 bias_ptr.add_offsets(c_idx),
649 None,
650 None,
651 &[],
652 None,
653 None,
654 None,
655 false,
656 ),
657 &[BLOCK_HW],
658 );
659
660 let batch_channel_offset: i32 = b * C * HW + c * HW;
662 let zeros = T::zeros::<D>(&[BLOCK_HW]);
663 let mut hw_start: i32 = 0;
664 while hw_start < HW {
665 let offsets = T::arange(0, BLOCK_HW) + hw_start;
666 let mask = offsets.lt(HW);
667 let elem_offsets = offsets + batch_channel_offset;
668
669 let x_tile = T::load(
670 x_ptr.add_offsets(elem_offsets),
671 Some(mask),
672 Some(zeros),
673 &[],
674 None,
675 None,
676 None,
677 false,
678 );
679 let y_tile = gamma * (x_tile - mean) * rstd + beta;
680 T::store(
681 y_ptr.add_offsets(elem_offsets),
682 y_tile,
683 Some(mask),
684 &[],
685 None,
686 None,
687 );
688
689 hw_start += BLOCK_HW;
690 }
691}
692
693pub struct BatchNorm2dNchwInferenceRuntimeOp<D: Float + Send + Sync + 'static> {
700 fwd: BatchNorm2dNchwForwardInference<D>,
701 block_hw: i32,
702 eps: f32,
703}
704
705impl<D: Float + Send + Sync + 'static> BatchNorm2dNchwInferenceRuntimeOp<D> {
706 pub fn new(block_hw: i32, eps: f32) -> Self {
707 Self {
708 fwd: BatchNorm2dNchwForwardInference::<D>::new(block_hw),
709 block_hw,
710 eps,
711 }
712 }
713
714 pub fn forward_source(&self) -> &str {
715 &self.fwd.source
716 }
717 pub fn kernel_name(&self) -> &str {
718 self.fwd.name
719 }
720
721 #[cfg(feature = "training")]
722 pub fn backward_source(&self) -> String {
723 BatchNorm2dNchwBackward::<D>::new(self.block_hw).source
724 }
725}
726
727impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp
728 for BatchNorm2dNchwInferenceRuntimeOp<D>
729{
730 fn n_activation_inputs(&self) -> usize {
731 1
732 }
733
734 fn param_shapes(&self, input_shapes: &[&[usize]], _output_shape: &[usize]) -> Vec<Vec<usize>> {
735 let c = input_shapes[0][1];
736 vec![vec![c], vec![c], vec![c], vec![c]]
737 }
738
739 fn param_names(&self) -> &'static [&'static str] {
740 &["weight", "bias", "running_mean", "running_var"]
741 }
742
743 fn pack_args(
744 &self,
745 inputs: &[(teeny_core::model::RawPtr, &[usize])],
746 params: &[teeny_core::model::RawPtr],
747 output: teeny_core::model::RawPtr,
748 output_shape: &[usize],
749 _output_row_stride: i32,
750 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
751 ) {
752 let c = output_shape[1] as i32;
753 let hw = (output_shape[2] * output_shape[3]) as i32;
754 visitor.visit_ptr(inputs[0].0);
755 visitor.visit_ptr(output);
756 visitor.visit_ptr(params[0]); visitor.visit_ptr(params[1]); visitor.visit_ptr(params[2]); visitor.visit_ptr(params[3]); visitor.visit_i32(c);
761 visitor.visit_i32(hw);
762 visitor.visit_f32(self.eps);
763 }
764
765 fn block(&self) -> [u32; 3] {
766 [self.block_hw as u32, 1, 1]
767 }
768
769 fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
770 [output_shape[1] as u32, output_shape[0] as u32, 1]
771 }
772
773 #[cfg(feature = "training")]
774 fn has_backward(&self) -> bool {
775 true
776 }
777
778 #[cfg(feature = "training")]
779 fn pack_backward_args(
780 &self,
781 inputs: &[(teeny_core::model::RawPtr, &[usize])],
782 params: &[teeny_core::model::RawPtr],
783 _output: teeny_core::model::RawPtr,
784 _output_shape: &[usize],
785 grad_output: teeny_core::model::RawPtr,
786 _grad_output_row_stride: i32,
787 grad_inputs: &[teeny_core::model::RawPtr],
788 grad_params: &[teeny_core::model::RawPtr],
789 visitor: &mut dyn teeny_core::device::program::ArgVisitor,
790 ) {
791 let in_shape = inputs[0].1; let b = in_shape[0] as i32;
794 let c = in_shape[1] as i32;
795 let hw = (in_shape[2] * in_shape[3]) as i32;
796 visitor.visit_ptr(grad_output); visitor.visit_ptr(inputs[0].0); visitor.visit_ptr(grad_inputs[0]); visitor.visit_ptr(params[0]); visitor.visit_ptr(params[2]); visitor.visit_ptr(params[3]); visitor.visit_ptr(grad_params[0]); visitor.visit_ptr(grad_params[1]); visitor.visit_i32(b);
805 visitor.visit_i32(c);
806 visitor.visit_i32(hw);
807 visitor.visit_f32(self.eps);
808 }
809
810 #[cfg(feature = "training")]
811 fn backward_block(&self) -> [u32; 3] {
812 [self.block_hw as u32, 1, 1]
813 }
814
815 #[cfg(feature = "training")]
816 fn backward_grid(&self, input_shapes: &[&[usize]], _output_shape: &[usize]) -> [u32; 3] {
817 [input_shapes[0][1] as u32, 1, 1]
819 }
820}
821
822#[cfg(feature = "training")]
838#[kernel]
839pub fn batch_norm_2d_nchw_backward<T: Triton, D: Float, const BLOCK_HW: i32>(
840 dy_ptr: T::Pointer<D>,
841 x_ptr: T::Pointer<D>,
842 dx_ptr: T::Pointer<D>,
843 weight_ptr: T::Pointer<D>,
844 running_mean_ptr: T::Pointer<D>,
845 running_var_ptr: T::Pointer<D>,
846 dweight_ptr: T::Pointer<D>,
847 dbias_ptr: T::Pointer<D>,
848 B: i32,
849 C: i32,
850 HW: i32,
851 eps: f32,
852) where
853 T::I32Tensor: types::Tensor<i32, 1>,
854 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
855 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
856{
857 let c = T::program_id(Axis::X);
858 let c_idx = T::arange(0, 1) + c;
859
860 let mean = T::load(
862 running_mean_ptr.add_offsets(c_idx),
863 None,
864 None,
865 &[],
866 None,
867 None,
868 None,
869 false,
870 );
871 let var = T::load(
872 running_var_ptr.add_offsets(c_idx),
873 None,
874 None,
875 &[],
876 None,
877 None,
878 None,
879 false,
880 );
881 let rstd = T::rsqrt(var + T::cast::<f32, D>(T::full::<f32>(&[1], eps), None, false));
882 let gamma = T::load(
883 weight_ptr.add_offsets(c_idx),
884 None,
885 None,
886 &[],
887 None,
888 None,
889 None,
890 false,
891 );
892
893 let mut sum_dy = T::zeros::<D>(&[1]);
898 let mut sum_dy_xhat = T::zeros::<D>(&[1]);
899 let total_bhw = B * HW;
900 let mut n: i32 = 0;
901 while n < total_bhw {
902 let b_idx = n / HW; let hw_idx = n % HW; let offset: i32 = b_idx * C * HW + c * HW + hw_idx;
905 let off_1 = T::arange(0, 1) + offset; let x_elem = T::load(
908 x_ptr.add_offsets(off_1),
909 None,
910 None,
911 &[],
912 None,
913 None,
914 None,
915 false,
916 );
917 let dy_elem = T::load(
918 dy_ptr.add_offsets(off_1),
919 None,
920 None,
921 &[],
922 None,
923 None,
924 None,
925 false,
926 );
927
928 let xhat = (x_elem - mean) * rstd;
929 sum_dy = sum_dy + dy_elem;
930 sum_dy_xhat = sum_dy_xhat + dy_elem * xhat;
931
932 let dx_elem = gamma * rstd * dy_elem;
934 T::store(dx_ptr.add_offsets(off_1), dx_elem, None, &[], None, None);
935
936 n += 1;
937 }
938
939 T::store(
940 dweight_ptr.add_offsets(c_idx),
941 sum_dy_xhat,
942 None,
943 &[],
944 None,
945 None,
946 );
947 T::store(dbias_ptr.add_offsets(c_idx), sum_dy, None, &[], None, None);
948}
949
950#[cfg(feature = "training")]
969#[kernel]
970pub fn batch_norm_backward<T: Triton, D: Float, const BLOCK_N: i32>(
971 dy_ptr: T::Pointer<D>,
972 x_ptr: T::Pointer<D>,
973 dx_ptr: T::Pointer<D>,
974 weight_ptr: T::Pointer<D>,
975 mean_ptr: T::Pointer<D>,
976 rstd_ptr: T::Pointer<D>,
977 dweight_ptr: T::Pointer<D>,
978 dbias_ptr: T::Pointer<D>,
979 N: i32,
980 C: i32,
981) where
982 T::I32Tensor: types::Tensor<i32, 1>,
983 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
984 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
985{
986 let c = T::program_id(Axis::X);
987 let c_idx = T::arange(0, 1) + c;
988
989 let mean = T::broadcast_to(
991 T::load(
992 mean_ptr.add_offsets(c_idx),
993 None,
994 None,
995 &[],
996 None,
997 None,
998 None,
999 false,
1000 ),
1001 &[BLOCK_N],
1002 );
1003 let rstd = T::broadcast_to(
1004 T::load(
1005 rstd_ptr.add_offsets(c_idx),
1006 None,
1007 None,
1008 &[],
1009 None,
1010 None,
1011 None,
1012 false,
1013 ),
1014 &[BLOCK_N],
1015 );
1016 let weight = T::broadcast_to(
1017 T::load(
1018 weight_ptr.add_offsets(c_idx),
1019 None,
1020 None,
1021 &[],
1022 None,
1023 None,
1024 None,
1025 false,
1026 ),
1027 &[BLOCK_N],
1028 );
1029
1030 let zeros_blk = T::zeros::<D>(&[BLOCK_N]);
1034 let mut acc_dy = zeros_blk;
1035 let mut acc_dy_xhat = zeros_blk;
1036 let mut n_start: i32 = 0;
1037
1038 while n_start < N {
1039 let offsets_n = T::arange(0, BLOCK_N) + n_start;
1040 let mask = offsets_n.lt(N);
1041 let elem_offsets = offsets_n * C + c;
1042
1043 let x_tile = T::load(
1044 x_ptr.add_offsets(elem_offsets),
1045 Some(mask),
1046 Some(zeros_blk),
1047 &[],
1048 None,
1049 None,
1050 None,
1051 false,
1052 );
1053 let dy_tile = T::load(
1054 dy_ptr.add_offsets(elem_offsets),
1055 Some(mask),
1056 Some(zeros_blk),
1057 &[],
1058 None,
1059 None,
1060 None,
1061 false,
1062 );
1063 let xhat = (x_tile - mean) * rstd;
1064
1065 acc_dy = acc_dy + dy_tile;
1066 acc_dy_xhat = acc_dy_xhat + dy_tile * xhat;
1067
1068 n_start += BLOCK_N;
1069 }
1070
1071 let sum_dy = T::sum(acc_dy, None, true);
1073 let sum_dy_xhat = T::sum(acc_dy_xhat, None, true);
1074
1075 T::store(
1077 dweight_ptr.add_offsets(c_idx),
1078 sum_dy_xhat,
1079 None,
1080 &[],
1081 None,
1082 None,
1083 );
1084 T::store(dbias_ptr.add_offsets(c_idx), sum_dy, None, &[], None, None);
1085
1086 let n_inv = T::broadcast_to(
1088 T::cast::<f32, D>(T::full::<f32>(&[1], 1.0f32 / (N as f32)), None, false),
1089 &[BLOCK_N],
1090 );
1091 let sum_dy_bcast = T::broadcast_to(sum_dy, &[BLOCK_N]);
1092 let sum_dy_xhat_bcast = T::broadcast_to(sum_dy_xhat, &[BLOCK_N]);
1093
1094 n_start = 0;
1096 while n_start < N {
1097 let offsets_n = T::arange(0, BLOCK_N) + n_start;
1098 let mask = offsets_n.lt(N);
1099 let elem_offsets = offsets_n * C + c;
1100
1101 let x_tile = T::load(
1102 x_ptr.add_offsets(elem_offsets),
1103 Some(mask),
1104 Some(zeros_blk),
1105 &[],
1106 None,
1107 None,
1108 None,
1109 false,
1110 );
1111 let dy_tile = T::load(
1112 dy_ptr.add_offsets(elem_offsets),
1113 Some(mask),
1114 Some(zeros_blk),
1115 &[],
1116 None,
1117 None,
1118 None,
1119 false,
1120 );
1121 let xhat = (x_tile - mean) * rstd;
1122
1123 let dx_tile =
1124 weight * rstd * (dy_tile - sum_dy_bcast * n_inv - xhat * sum_dy_xhat_bcast * n_inv);
1125
1126 T::store(
1127 dx_ptr.add_offsets(elem_offsets),
1128 dx_tile,
1129 Some(mask),
1130 &[],
1131 None,
1132 None,
1133 );
1134
1135 n_start += BLOCK_N;
1136 }
1137}