Skip to main content

combs_models/
norm.rs

1//! RMSNorm (Llama-style, no learnable bias).
2
3use burn::tensor::{Tensor, backend::Backend};
4
5/// `y = x / rms(x) * w` where `rms` is taken over the last dimension and
6/// `eps` is added inside the square root.
7pub fn rms_norm<B: Backend, const D: usize>(
8    x: Tensor<B, D>,
9    weight: Tensor<B, 1>,
10    eps: f64,
11) -> Tensor<B, D> {
12    let dims = x.dims();
13    let hidden = dims[D - 1];
14
15    // mean(x^2) over the last dim, keeping rank for broadcasting.
16    let mean_sq = x.clone().powf_scalar(2.0).mean_dim(D - 1);
17    let inv_rms = mean_sq.add_scalar(eps).sqrt().recip();
18
19    let mut shape = [1usize; D];
20    shape[D - 1] = hidden;
21    x * inv_rms * weight.reshape(shape)
22}
23
24/// LayerNorm (SigLIP-style, learnable weight + bias):
25/// `y = (x - μ) / sqrt(σ² + eps) * w + b`, statistics over the last dim.
26pub fn layer_norm<B: Backend, const D: usize>(
27    x: Tensor<B, D>,
28    weight: Tensor<B, 1>,
29    bias: Tensor<B, 1>,
30    eps: f64,
31) -> Tensor<B, D> {
32    let dims = x.dims();
33    let hidden = dims[D - 1];
34
35    let mean = x.clone().mean_dim(D - 1);
36    let centered = x - mean;
37    let var = centered.clone().powf_scalar(2.0).mean_dim(D - 1);
38    let inv_std = var.add_scalar(eps).sqrt().recip();
39
40    let mut shape = [1usize; D];
41    shape[D - 1] = hidden;
42    centered * inv_std * weight.reshape(shape) + bias.reshape(shape)
43}
44
45#[cfg(test)]
46mod tests {
47    use super::*;
48    use burn::tensor::TensorData;
49
50    type TestBackend = burn::backend::NdArray<f32>;
51
52    #[test]
53    fn normalizes_rows_to_unit_rms() {
54        let device = burn::tensor::Device::<TestBackend>::default();
55        let x: Tensor<TestBackend, 2> = Tensor::from_data(
56            TensorData::new(vec![3.0f32, 4.0, 1.0, -2.0, 0.5, 2.5], [2, 3]),
57            &device,
58        );
59        let w: Tensor<TestBackend, 1> = Tensor::ones([3], &device);
60        let y = rms_norm(x, w, 1e-6);
61        // With unit weights, each row of y must have RMS == 1 (up to eps).
62        let rms = y
63            .powf_scalar(2.0)
64            .mean_dim(1)
65            .sqrt()
66            .into_data()
67            .to_vec::<f32>()
68            .unwrap();
69        for (i, r) in rms.iter().enumerate() {
70            assert!((r - 1.0).abs() < 1e-4, "row {i} rms = {r}");
71        }
72    }
73
74    #[test]
75    fn applies_weight() {
76        let device = burn::tensor::Device::<TestBackend>::default();
77        let x: Tensor<TestBackend, 2> =
78            Tensor::from_data(TensorData::new(vec![1.0f32, 2.0], [1, 2]), &device);
79        let w: Tensor<TestBackend, 1> =
80            Tensor::from_data(TensorData::new(vec![2.0f32, 2.0], [2]), &device);
81        let y = rms_norm(x.clone(), w, 1e-6);
82        let z = rms_norm(x, Tensor::ones([2], &device), 1e-6);
83        let yv: Vec<f32> = y.into_data().to_vec().unwrap();
84        let zv: Vec<f32> = z.into_data().to_vec().unwrap();
85        for i in 0..2 {
86            assert!((yv[i] - 2.0 * zv[i]).abs() < 1e-4);
87        }
88    }
89}