Skip to main content

teeny_kernels/nn/norm/
rmsnorm.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//! RMSNorm Triton kernels.
18//!
19//! RMSNorm normalises each row by its root-mean-square (no mean subtraction):
20//!   rms[m]    = sqrt( (1/N) * Σ_n x[m,n]² + eps )
21//!   y[m,n]    = x[m,n] / rms[m] * γ[n]
22//!
23//! Grid: `[M]` — one CTA per row. Layout identical to LayerNorm.
24
25#![allow(non_snake_case)]
26
27use teeny_core::dtype::Float;
28use teeny_macros::kernel;
29use teeny_triton::triton::{
30    types::{AddOffsets, Comparison},
31    *,
32};
33
34// ─── Forward ─────────────────────────────────────────────────────────────────
35
36/// RMSNorm forward pass.
37///
38/// Grid: `[M]` — one CTA per row.
39#[kernel]
40pub fn rms_norm_forward<T: Triton, D: Float, const BLOCK_N: i32>(
41    x_ptr: T::Pointer<D>,
42    y_ptr: T::Pointer<D>,
43    weight_ptr: T::Pointer<D>,
44    rrms_ptr: T::Pointer<D>,
45    _M: i32,
46    N: i32,
47    eps: f32,
48) where
49    T::I32Tensor: types::Tensor<i32, 1>,
50    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
51    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
52{
53    let row = T::program_id(Axis::X);
54    let row_start = row * N;
55    let row_idx = T::arange(0, 1) + row;
56
57    let zeros = T::zeros::<D>(&[BLOCK_N]);
58    let zero_1 = T::zeros::<D>(&[1]);
59    let n_inv = T::cast::<f32, D>(T::full::<f32>(&[1], 1.0f32 / (N as f32)), None, false);
60
61    // ── Pass 1: accumulate Σ x² ──────────────────────────────────────────────
62    let mut sq_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        sq_sum = sq_sum + T::sum(x_tile * x_tile, None, true);
78        n_start += BLOCK_N;
79    }
80    let eps_t = T::cast::<f32, D>(T::full::<f32>(&[1], eps), None, false);
81    let rrms_1 = T::rsqrt(sq_sum * n_inv + eps_t);
82    let rrms = T::broadcast_to(rrms_1, &[BLOCK_N]);
83
84    T::store(rrms_ptr.add_offsets(row_idx), rrms_1, None, &[], None, None);
85
86    // ── Pass 2: normalise ─────────────────────────────────────────────────────
87    n_start = 0;
88    while n_start < N {
89        let col_offs = T::arange(0, BLOCK_N) + n_start;
90        let mask = col_offs.lt(N);
91        let x_tile = T::load(
92            x_ptr.add_offsets(col_offs + row_start),
93            Some(mask),
94            Some(zeros),
95            &[],
96            None,
97            None,
98            None,
99            false,
100        );
101        let gamma = T::load(
102            weight_ptr.add_offsets(col_offs),
103            Some(mask),
104            Some(zeros),
105            &[],
106            None,
107            None,
108            None,
109            false,
110        );
111        let y_tile = x_tile * rrms * gamma;
112        T::store(
113            y_ptr.add_offsets(col_offs + row_start),
114            y_tile,
115            Some(mask),
116            &[],
117            None,
118            None,
119        );
120        n_start += BLOCK_N;
121    }
122}
123
124// ─── Backward ────────────────────────────────────────────────────────────────
125
126/// RMSNorm backward pass.
127///
128/// ```text
129/// dx[m,n] = rrms[m] * γ[n] * (dy[m,n] - x[m,n] * rrms[m]² * Σ_n dy[m,n]*γ[n]*x[m,n] / N)
130/// dweight[n] = Σ_m dy[m,n] * x[m,n] * rrms[m]
131/// ```
132///
133/// Grid: `[M]` — one CTA per row.
134#[cfg(feature = "training")]
135#[kernel]
136pub fn rms_norm_backward<T: Triton, D: Float, const BLOCK_N: i32>(
137    dy_ptr: T::Pointer<D>,
138    x_ptr: T::Pointer<D>,
139    dx_ptr: T::Pointer<D>,
140    weight_ptr: T::Pointer<D>,
141    dweight_ptr: T::Pointer<D>,
142    rrms_ptr: T::Pointer<D>,
143    _M: i32,
144    N: i32,
145) where
146    T::I32Tensor: types::Tensor<i32, 1>,
147    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
148    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
149{
150    let row = T::program_id(Axis::X);
151    let row_start = row * N;
152    let row_idx = T::arange(0, 1) + row;
153
154    let zeros = T::zeros::<D>(&[BLOCK_N]);
155    let zero_1 = T::zeros::<D>(&[1]);
156    let n_inv = T::cast::<f32, D>(T::full::<f32>(&[1], 1.0f32 / (N as f32)), None, false);
157
158    let rrms_1 = T::load(
159        rrms_ptr.add_offsets(row_idx),
160        None,
161        None,
162        &[],
163        None,
164        None,
165        None,
166        false,
167    );
168    let rrms = T::broadcast_to(rrms_1, &[BLOCK_N]);
169
170    // ── Pass 1: Σ dy * γ * x ─────────────────────────────────────────────────
171    let mut dot = zero_1;
172    let mut n_start: i32 = 0;
173    while n_start < N {
174        let col_offs = T::arange(0, BLOCK_N) + n_start;
175        let mask = col_offs.lt(N);
176        let x_tile = T::load(
177            x_ptr.add_offsets(col_offs + row_start),
178            Some(mask),
179            Some(zeros),
180            &[],
181            None,
182            None,
183            None,
184            false,
185        );
186        let dy_tile = T::load(
187            dy_ptr.add_offsets(col_offs + row_start),
188            Some(mask),
189            Some(zeros),
190            &[],
191            None,
192            None,
193            None,
194            false,
195        );
196        let gamma = T::load(
197            weight_ptr.add_offsets(col_offs),
198            Some(mask),
199            Some(zeros),
200            &[],
201            None,
202            None,
203            None,
204            false,
205        );
206        dot = dot + T::sum(dy_tile * gamma * x_tile, None, true);
207        n_start += BLOCK_N;
208    }
209    let rrms_sq = T::broadcast_to(rrms_1 * rrms_1, &[BLOCK_N]);
210    let scale = T::broadcast_to(dot * n_inv, &[BLOCK_N]);
211
212    // ── Pass 2: dx and dweight ────────────────────────────────────────────────
213    n_start = 0;
214    while n_start < N {
215        let col_offs = T::arange(0, BLOCK_N) + n_start;
216        let mask = col_offs.lt(N);
217        let x_tile = T::load(
218            x_ptr.add_offsets(col_offs + row_start),
219            Some(mask),
220            Some(zeros),
221            &[],
222            None,
223            None,
224            None,
225            false,
226        );
227        let dy_tile = T::load(
228            dy_ptr.add_offsets(col_offs + row_start),
229            Some(mask),
230            Some(zeros),
231            &[],
232            None,
233            None,
234            None,
235            false,
236        );
237        let gamma = T::load(
238            weight_ptr.add_offsets(col_offs),
239            Some(mask),
240            Some(zeros),
241            &[],
242            None,
243            None,
244            None,
245            false,
246        );
247        let dw_old = T::load(
248            dweight_ptr.add_offsets(col_offs),
249            Some(mask),
250            Some(zeros),
251            &[],
252            None,
253            None,
254            None,
255            false,
256        );
257
258        let dx_tile = rrms * gamma * (dy_tile - x_tile * rrms_sq * scale);
259        T::store(
260            dx_ptr.add_offsets(col_offs + row_start),
261            dx_tile,
262            Some(mask),
263            &[],
264            None,
265            None,
266        );
267        T::store(
268            dweight_ptr.add_offsets(col_offs),
269            dw_old + dy_tile * x_tile * rrms,
270            Some(mask),
271            &[],
272            None,
273            None,
274        );
275        n_start += BLOCK_N;
276    }
277}