Skip to main content

teeny_kernels/nn/norm/
batchnorm.rs

1/*
2 * Copyright (c) 2026 Teenygrad.
3 *
4 * Licensed under the Apache License, Version 2.0 (the "License");
5 * you may not use this file except in compliance with the License.
6 * You may obtain a copy of the License at
7 *
8 *   http://www.apache.org/licenses/LICENSE-2.0
9 *
10 * Unless required by applicable law or agreed to in writing, software
11 * distributed under the License is distributed on an "AS IS" BASIS,
12 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13 * See the License for the specific language governing permissions and
14 * limitations under the License.
15 */
16
17//! BatchNorm1d Triton kernels.
18//!
19//! Layout: input `x` is `[N, C]` row-major. Element `x[n, c]` lives at flat
20//! offset `n * C + c`.
21//!
22//! Parallelism: **one CTA per channel**. Each CTA iterates over all N batch
23//! elements in `BLOCK_N`-wide tiles. This avoids cross-CTA synchronisation
24//! entirely — C channels execute concurrently across SMs.
25//!
26//! Training requires two sequential kernel launches separated by a host sync:
27//!   1. `batch_norm_stats_forward`   — computes per-channel mean + rstd, updates
28//!      running stats.
29//!   2. `batch_norm_normalize_forward` — normalises x using the saved stats.
30//!
31//! Inference uses a single kernel that reads the frozen running statistics.
32
33#![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// ─── Inference: single kernel, frozen running statistics ─────────────────────
43
44/// Normalises input `x` using the frozen `running_mean` / `running_var`.
45///
46/// Grid: `[C]` — one CTA per channel.
47#[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    // Load per-channel scalars (shape [1]) and broadcast to [BLOCK_N].
67    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    // Normalise all N elements for this channel.
122    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// ─── Training: kernel 1 — compute per-channel statistics ─────────────────────
155
156/// Computes per-channel mean and rstd from the current mini-batch, saves them
157/// for the normalisation kernel and the backward pass, and updates the running
158/// statistics with exponential moving average.
159///
160/// Grid: `[C]` — one CTA per channel.
161///
162/// **Must complete (host sync) before `batch_norm_normalize_forward` is launched.**
163#[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    // Accumulate sum(x) and sum(x²) over all N elements for this channel.
183    // Triton idiom: accumulate BLOCK_N-wide tiles inside the loop, reduce once
184    // outside — tt.reduce inside a loop body is not supported by Triton's lowering.
185    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    // Single reduce outside the loop — shape [BLOCK_N] → [1].
212    let sum = T::sum(acc_sum, None, true);
213    let sum_sq = T::sum(acc_sum_sq, None, true);
214
215    // Derive mean, biased variance, and rstd (all shape [1]).
216    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    // Save for the normalisation and backward kernels.
222    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    // Exponential moving average: running = (1 - m) * running + m * batch.
227    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// ─── Training: kernel 2 — normalise using saved statistics ───────────────────
269
270/// Normalises x using the mean and rstd produced by `batch_norm_stats_forward`.
271///
272/// Grid: `[C]` — one CTA per channel.
273#[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    // Load per-channel scalars and broadcast to [BLOCK_N].
293    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// ─── Training RuntimeOp implementations ──────────────────────────────────────
379
380/// RuntimeOp for the stats kernel node in a training BatchNorm graph.
381///
382/// Stores `eps` and `momentum` (not in the macro-generated struct) and handles
383/// the packed `[2*C]` output layout: first C elements = mean, last C = rstd.
384#[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/// RuntimeOp for the normalize kernel node in a training BatchNorm graph.
453///
454/// Expects two activation inputs: `inputs[0]` = x, `inputs[1]` = packed stats
455/// `[2*C]` from the stats node (first C = mean, last C = rstd).
456#[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        // kernel args: dy, x, dx, weight, mean, rstd, dweight, dbias, N, C
545        // inputs[0] = x,     inputs[1] = stats [2*C] (mean at [0..C], rstd at [C..2C])
546        // params[0] = weight, params[1] = bias (unused for dx)
547        // grad_inputs[0] = dx,  grad_params[0] = dweight,  grad_params[1] = dbias
548        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); // dy_ptr
554        visitor.visit_ptr(inputs[0].0); // x_ptr
555        visitor.visit_ptr(grad_inputs[0]); // dx_ptr
556        visitor.visit_ptr(params[0]); // weight_ptr
557        visitor.visit_ptr(mean_ptr); // mean_ptr
558        visitor.visit_ptr(rstd_ptr); // rstd_ptr
559        visitor.visit_ptr(grad_params[0]); // dweight_ptr
560        visitor.visit_ptr(grad_params[1]); // dbias_ptr
561        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        // input_shapes[1] = [2*C] (stats); one CTA per channel
571        let c = input_shapes[1][0] / 2;
572        [c as u32, 1, 1]
573    }
574}
575
576// ─── Inference (NCHW): NCHW-native BatchNorm2d ───────────────────────────────
577
578/// Normalises NCHW input `x` using frozen running statistics.
579///
580/// Input layout: [B, C, H, W] row-major. Element `x[b, c, h, w]` lives at
581/// offset `b*C*HW + c*HW + h*W + w`.
582///
583/// Grid: `[C, B]` — one CTA per (channel, batch) pair; each CTA iterates over
584/// H*W spatial positions in `BLOCK_HW`-wide tiles.
585#[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    // Load per-channel scalars and broadcast to [BLOCK_HW].
606    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    // Flat start offset for (b, c, hw=0) in NCHW: b*C*HW + c*HW
661    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
693// ─── Inference (NCHW) RuntimeOp ──────────────────────────────────────────────
694
695/// RuntimeOp for NCHW BatchNorm2d inference.
696///
697/// Parameter layout (4 params): `[weight, bias, running_mean, running_var]`,
698/// each of shape `[C]`.
699pub 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]); // weight
757        visitor.visit_ptr(params[1]); // bias
758        visitor.visit_ptr(params[2]); // running_mean
759        visitor.visit_ptr(params[3]); // running_var
760        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        // kernel args: dy, x, dx, weight, running_mean, running_var, dweight, dbias, B, C, HW, eps
792        let in_shape = inputs[0].1; // [B, C, H, W]
793        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); // dy_ptr
797        visitor.visit_ptr(inputs[0].0); // x_ptr
798        visitor.visit_ptr(grad_inputs[0]); // dx_ptr
799        visitor.visit_ptr(params[0]); // weight_ptr
800        visitor.visit_ptr(params[2]); // running_mean_ptr
801        visitor.visit_ptr(params[3]); // running_var_ptr
802        visitor.visit_ptr(grad_params[0]); // dweight_ptr
803        visitor.visit_ptr(grad_params[1]); // dbias_ptr
804        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] = [B, C, H, W]; one CTA per channel
818        [input_shapes[0][1] as u32, 1, 1]
819    }
820}
821
822// ─── Training (NCHW): backward pass ──────────────────────────────────────────
823
824/// Computes gradients for NCHW BatchNorm2d inference.
825///
826/// Since `running_mean` and `running_var` are frozen constants, the backward of
827/// `y = gamma * (x - mean) * rstd + beta` (with respect to `x`) is simply:
828/// ```text
829/// dx[b,c,h,w]   = gamma[c] * rstd[c] * dy[b,c,h,w]
830/// dweight[c]     = Σ_{b,h,w} dy * xhat
831/// dbias[c]       = Σ_{b,h,w} dy
832/// ```
833///
834/// A single loop over (b, hw) computes all three simultaneously.
835///
836/// Grid: `[C]` — one CTA per channel.
837#[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    // Load per-channel scalars as [1]-shaped tensors to match element-wise loop.
861    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    // Use a single flat loop over B*HW with one element per iteration.
894    // Scalar b_idx/hw_idx arithmetic avoids tensor-level division/modulo,
895    // and a single-level loop with [1] accumulators is correctly lowered
896    // (same pattern as batch_norm_stats_forward).
897    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; // scalar division — always valid
903        let hw_idx = n % HW; // scalar modulo
904        let offset: i32 = b_idx * C * HW + c * HW + hw_idx;
905        let off_1 = T::arange(0, 1) + offset; // [1] pointing at this element
906
907        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        // dx = gamma * rstd * dy (frozen-stats: mean/rstd are constants)
933        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// ─── Training (NC): backward pass ─────────────────────────────────────────────
951
952/// Computes gradients for BatchNorm.
953///
954/// Given saved `mean` and `rstd` from the forward pass:
955/// ```text
956/// xhat      = (x - mean) * rstd
957/// dbias[c]  = Σ_n dy[n,c]
958/// dweight[c]= Σ_n dy[n,c] * xhat[n,c]
959/// dx[n,c]   = weight[c] * rstd[c] * (dy[n,c]
960///               - dbias[c] / N
961///               - xhat[n,c] * dweight[c] / N)
962/// ```
963///
964/// Uses two sequential passes over N within the same CTA to avoid storing
965/// the full xhat tensor.
966///
967/// Grid: `[C]` — one CTA per channel.
968#[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    // Load per-channel scalars; broadcast to [BLOCK_N] for element-wise ops.
990    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    // Pass 1: accumulate dbias (= Σ dy) and dweight (= Σ dy * xhat).
1031    // Triton idiom: accumulate BLOCK_N-wide tiles inside the loop, reduce once
1032    // outside — tt.reduce inside a loop body is not supported by Triton's lowering.
1033    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    // Single reduce outside the loop — shape [BLOCK_N] → [1].
1072    let sum_dy = T::sum(acc_dy, None, true);
1073    let sum_dy_xhat = T::sum(acc_dy_xhat, None, true);
1074
1075    // Save dweight and dbias (shape [1] → stored as scalars).
1076    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    // Broadcast reduction results and 1/N for pass 2.
1087    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    // Pass 2: compute dx.
1095    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}