1use burn::tensor::{Tensor, backend::Backend};
4
5pub 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 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
24pub 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 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}