Skip to main content

teeny_kernels/nn/norm/
layernorm.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//! LayerNorm Triton kernels.
18//!
19//! Layout: input `x` is `[M, N]` row-major where M = product of batch / outer
20//! dimensions and N = product of the normalized dimensions.  Each CTA handles
21//! one row (one sample), reading all N elements in `BLOCK_N`-wide tiles.
22//!
23//! Forward: y[m, n] = (x[m, n] − mean_m) / sqrt(var_m + eps) * γ[n] + β[n]
24//!
25//! Training launches a single forward kernel that also writes out the saved
26//! `mean` and `rstd` buffers for the backward pass.
27
28#![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// ─── Inference ───────────────────────────────────────────────────────────────
38
39/// Forward pass using pre-computed running statistics (inference only).
40///
41/// Grid: `[M]` — one CTA per row.
42#[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    // ── Pass 1: accumulate mean ───────────────────────────────────────────────
60    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    // ── Pass 2: accumulate variance ───────────────────────────────────────────
85    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        // Mask the diff so out-of-bounds positions don't contribute mean^2 to variance.
101        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    // ── Pass 3: normalise and apply affine transform ──────────────────────────
109    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// ─── Training forward ─────────────────────────────────────────────────────────
157
158/// Forward pass that also saves per-row mean and rstd for the backward pass.
159///
160/// Grid: `[M]` — one CTA per row.
161#[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    // ── Pass 1: mean ─────────────────────────────────────────────────────────
187    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    // ── Pass 2: variance ─────────────────────────────────────────────────────
209    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        // Mask the diff so out-of-bounds positions don't contribute mean^2 to variance.
225        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    // ── Pass 3: normalise ─────────────────────────────────────────────────────
237    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// ─── Training backward ───────────────────────────────────────────────────────
285
286/// Backward pass for LayerNorm.
287///
288/// Given saved `mean` and `rstd` from the forward pass:
289/// ```text
290/// xhat[m,n]    = (x[m,n] - mean[m]) * rstd[m]
291/// dweight[n]   = Σ_m dy[m,n] * xhat[m,n]
292/// dbias[n]     = Σ_m dy[m,n]
293/// dx[m,n]      = rstd[m] * γ[n] * (dy[m,n]
294///                  - (Σ_n dy[m,n]*γ[n]) / N
295///                  - xhat[m,n] * (Σ_n dy[m,n]*γ[n]*xhat[m,n]) / N)
296/// ```
297///
298/// Grid: `[M]` — one CTA per row.
299#[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    // ── Pass 1: accumulate row-level dot products ─────────────────────────────
349    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    // ── Pass 2: compute dx and accumulate dweight / dbias ────────────────────
394    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
480// ─── Inference RuntimeOp ──────────────────────────────────────────────────────
481
482/// RuntimeOp for LayerNorm inference.
483///
484/// Parameter layout (2 params): `[weight, bias]`, each of shape `[N]` where
485/// N is the last (normalized) dimension of the input.
486pub 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        // N = last dim of input
519        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); // x
542        visitor.visit_ptr(output); // y
543        visitor.visit_ptr(params[0]); // weight (gamma)
544        visitor.visit_ptr(params[1]); // bias (beta)
545        visitor.visit_i32(m); // M
546        visitor.visit_i32(n); // N
547        visitor.visit_f32(self.eps); // eps
548    }
549
550    fn block(&self) -> [u32; 3] {
551        [1, 1, 1]
552    }
553
554    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
555        // one CTA per row; M = product of all dims except last
556        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}