Skip to main content

zsfm_nn/
norm.rs

1use anyhow::Result;
2use candle_core::{Tensor, D};
3
4/// RMSNorm over the last dim: `x / sqrt(mean(x^2) + eps) * weight` (weight optional — some
5/// models fold the scale into a separate op and call this with `None`). Computes in `x`'s own
6/// dtype — cast beforehand if a caller needs a fixed compute precision regardless of input dtype.
7pub fn rms_norm(x: &Tensor, weight: Option<&Tensor>, eps: f64) -> Result<Tensor> {
8    let rms = x.sqr()?.mean_keepdim(D::Minus1)?;
9    let rms = (rms + eps)?.sqrt()?;
10    let x = x.broadcast_div(&rms)?;
11    match weight {
12        Some(w) => Ok(x.broadcast_mul(w)?),
13        None => Ok(x),
14    }
15}
16
17/// Standard LayerNorm over the last dim: `(x - mean) / sqrt(var + eps) * weight + bias`.
18pub fn layer_norm(x: &Tensor, weight: &Tensor, bias: &Tensor, eps: f64) -> Result<Tensor> {
19    let mean = x.mean_keepdim(D::Minus1)?;
20    let x = x.broadcast_sub(&mean)?;
21    let var = x.sqr()?.mean_keepdim(D::Minus1)?;
22    let std = (var + eps)?.sqrt()?;
23    let x = x.broadcast_div(&std)?;
24    let x = x.broadcast_mul(weight)?;
25    Ok(x.broadcast_add(bias)?)
26}
27
28#[cfg(test)]
29mod tests {
30    use super::*;
31    use candle_core::Device;
32
33    #[test]
34    fn rms_norm_matches_manual_formula() {
35        let device = Device::Cpu;
36        let x = Tensor::from_vec(vec![1.0f32, 2.0, 3.0, -4.0], (1, 4), &device).unwrap();
37        let w = Tensor::from_vec(vec![2.0f32, 2.0, 2.0, 2.0], (4,), &device).unwrap();
38        let eps = 1e-6;
39
40        let got: Vec<f32> = rms_norm(&x, Some(&w), eps)
41            .unwrap()
42            .flatten_all()
43            .unwrap()
44            .to_vec1()
45            .unwrap();
46
47        let vals = [1.0f32, 2.0, 3.0, -4.0];
48        let mean_sq: f32 = vals.iter().map(|v| v * v).sum::<f32>() / 4.0;
49        let rms = (mean_sq + eps as f32).sqrt();
50        let want: Vec<f32> = vals.iter().map(|v| v / rms * 2.0).collect();
51
52        for (g, w) in got.iter().zip(want.iter()) {
53            assert!((g - w).abs() < 1e-5, "{g} vs {w}");
54        }
55    }
56
57    #[test]
58    fn rms_norm_without_weight_is_identity_scaled() {
59        let device = Device::Cpu;
60        let x = Tensor::from_vec(vec![3.0f32, 4.0], (1, 2), &device).unwrap();
61        let got: Vec<f32> = rms_norm(&x, None, 0.0)
62            .unwrap()
63            .flatten_all()
64            .unwrap()
65            .to_vec1()
66            .unwrap();
67        // rms = sqrt((9+16)/2) = sqrt(12.5)
68        let rms = 12.5f32.sqrt();
69        assert!((got[0] - 3.0 / rms).abs() < 1e-6);
70        assert!((got[1] - 4.0 / rms).abs() < 1e-6);
71    }
72
73    #[test]
74    fn layer_norm_matches_manual_formula() {
75        let device = Device::Cpu;
76        let x = Tensor::from_vec(vec![1.0f32, 2.0, 3.0, 4.0], (1, 4), &device).unwrap();
77        let w = Tensor::from_vec(vec![1.0f32; 4], (4,), &device).unwrap();
78        let b = Tensor::from_vec(vec![0.0f32; 4], (4,), &device).unwrap();
79        let eps = 1e-5;
80
81        let got: Vec<f32> = layer_norm(&x, &w, &b, eps)
82            .unwrap()
83            .flatten_all()
84            .unwrap()
85            .to_vec1()
86            .unwrap();
87
88        let vals = [1.0f32, 2.0, 3.0, 4.0];
89        let mean = vals.iter().sum::<f32>() / 4.0;
90        let var = vals.iter().map(|v| (v - mean).powi(2)).sum::<f32>() / 4.0;
91        let std = (var + eps as f32).sqrt();
92        let want: Vec<f32> = vals.iter().map(|v| (v - mean) / std).collect();
93
94        for (g, w) in got.iter().zip(want.iter()) {
95            assert!((g - w).abs() < 1e-5, "{g} vs {w}");
96        }
97    }
98}