1use anyhow::Result;
2use candle_core::{Tensor, D};
3
4pub 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
17pub 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 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}