burn_nn/modules/norm/
layer.rs

1use burn_core as burn;
2
3use burn::config::Config;
4use burn::module::Content;
5use burn::module::DisplaySettings;
6use burn::module::Initializer;
7use burn::module::Module;
8use burn::module::ModuleDisplay;
9use burn::module::Param;
10use burn::tensor::Tensor;
11use burn::tensor::backend::Backend;
12
13/// Configuration to create a [LayerNorm](LayerNorm) layer using the [init function](LayerNormConfig::init).
14#[derive(Debug, Config)]
15pub struct LayerNormConfig {
16    /// The size of the input features.
17    pub d_model: usize,
18    /// A value required for numerical stability. Default: 1e-5
19    #[config(default = 1e-5)]
20    pub epsilon: f64,
21}
22
23/// Applies Layer Normalization over an input tensor as described in the paper [Layer Normalization](https://arxiv.org/abs/1607.06450).
24///
25/// `Y = norm(X) * γ + β`
26///
27/// Where:
28/// - `X` is the input tensor
29/// - `Y` is the output tensor
30/// - `γ` is the learnable weight
31/// - `β` is the learnable bias
32///
33/// Should be created using [LayerNormConfig](LayerNormConfig).
34#[derive(Module, Debug)]
35#[module(custom_display)]
36pub struct LayerNorm<B: Backend> {
37    /// The learnable weight.
38    pub gamma: Param<Tensor<B, 1>>,
39    /// The learnable bias.
40    pub beta: Param<Tensor<B, 1>>,
41    /// A value required for numerical stability.
42    epsilon: f64,
43}
44
45impl LayerNormConfig {
46    /// Initialize a new [layer norm](LayerNorm) module.
47    pub fn init<B: Backend>(&self, device: &B::Device) -> LayerNorm<B> {
48        let gamma = Initializer::Ones.init([self.d_model], device);
49        let beta = Initializer::Zeros.init([self.d_model], device);
50
51        LayerNorm {
52            gamma,
53            beta,
54            epsilon: self.epsilon,
55        }
56    }
57}
58
59impl<B: Backend> LayerNorm<B> {
60    /// Applies the forward pass on the input tensor.
61    ///
62    /// See the [LayerNorm](LayerNorm) documentation for more information.
63    ///
64    /// # Shapes
65    ///
66    /// - input: `[..., any, d_model]`
67    /// - output: `[..., any, d_model]`
68    pub fn forward<const D: usize>(&self, input: Tensor<B, D>) -> Tensor<B, D> {
69        let (var, mean) = input.clone().var_mean_bias(D - 1);
70
71        let input_normalized = input.sub(mean).div(var.add_scalar(self.epsilon).sqrt());
72
73        input_normalized
74            .mul(self.gamma.val().unsqueeze())
75            .add(self.beta.val().unsqueeze())
76    }
77}
78
79impl<B: Backend> ModuleDisplay for LayerNorm<B> {
80    fn custom_settings(&self) -> Option<DisplaySettings> {
81        DisplaySettings::new()
82            .with_new_line_after_attribute(false)
83            .optional()
84    }
85
86    fn custom_content(&self, content: Content) -> Option<Content> {
87        let [d_model] = self.gamma.shape().dims();
88        content
89            .add("d_model", &d_model)
90            .add("epsilon", &self.epsilon)
91            .optional()
92    }
93}
94
95#[cfg(test)]
96mod tests {
97    use super::*;
98    use alloc::format;
99    use burn::tensor::TensorData;
100    use burn::tensor::{Tolerance, ops::FloatElem};
101    type FT = FloatElem<TestBackend>;
102
103    #[cfg(feature = "std")]
104    use crate::{TestAutodiffBackend, TestBackend};
105
106    #[cfg(not(feature = "std"))]
107    use crate::TestBackend;
108
109    #[test]
110    fn layer_norm_forward() {
111        let device = Default::default();
112        let module = LayerNormConfig::new(10).init::<TestBackend>(&device);
113        let input = Tensor::<TestBackend, 2>::from_data(
114            TensorData::from([[
115                -0.6897, -2.7106, 2.2222, -1.0330, -0.8933, 1.1765, 0.0601, 1.5252, -0.3630, 0.6728,
116            ]]),
117            &device,
118        );
119
120        let output = module.forward(input);
121
122        let expected = TensorData::from([[
123            -0.4990, -1.9680, 1.6178, -0.7486, -0.6470, 0.8576, 0.0461, 1.1111, -0.2614, 0.4915,
124        ]]);
125        output
126            .to_data()
127            .assert_approx_eq::<FT>(&expected, Tolerance::default());
128    }
129
130    #[test]
131    fn layer_norm_forward_large_epsilon() {
132        let device = Default::default();
133        let module = LayerNormConfig::new(10)
134            .with_epsilon(1e-1)
135            .init::<TestBackend>(&device);
136        let input = Tensor::<TestBackend, 2>::from_data(
137            TensorData::from([[
138                -0.6897, -2.7106, 2.2222, -1.0330, -0.8933, 1.1765, 0.0601, 1.5252, -0.3630, 0.6728,
139            ]]),
140            &device,
141        );
142
143        let output = module.forward(input);
144
145        let expected = TensorData::from([[
146            -0.4863, -1.9180, 1.5766, -0.7295, -0.6305, 0.8358, 0.0449, 1.0828, -0.2548, 0.4790,
147        ]]);
148        output
149            .to_data()
150            .assert_approx_eq::<FT>(&expected, Tolerance::default());
151    }
152
153    #[cfg(feature = "std")]
154    #[test]
155    fn layer_norm_backward() {
156        let device = Default::default();
157        let module = LayerNormConfig::new(2).init::<TestAutodiffBackend>(&device);
158        let tensor_1 = Tensor::<TestAutodiffBackend, 2>::from_data(
159            TensorData::from([[0.0, 1.0], [3.0, 4.0]]),
160            &device,
161        )
162        .require_grad();
163        let tensor_2 = Tensor::<TestAutodiffBackend, 2>::from_data(
164            TensorData::from([[6.0, 7.0], [9.0, 10.0]]),
165            &device,
166        )
167        .require_grad();
168
169        let x = tensor_1.clone().matmul(tensor_2.clone());
170
171        let output = module.forward(x);
172        let grads = output.backward();
173
174        let tensor_1_grad = tensor_1.grad(&grads).unwrap();
175        let tensor_2_grad = tensor_2.grad(&grads).unwrap();
176        let gamma_grad = module.gamma.grad(&grads).unwrap();
177        let beta_grad = module.beta.grad(&grads).unwrap();
178
179        let expected = TensorData::from([-2.0, 2.0]);
180        gamma_grad
181            .to_data()
182            .assert_approx_eq::<FT>(&expected, Tolerance::default());
183
184        let expected = TensorData::from([2.0, 2.0]);
185        beta_grad
186            .to_data()
187            .assert_approx_eq::<FT>(&expected, Tolerance::default());
188
189        let expected = TensorData::zeros::<f32, _>(tensor_1_grad.shape());
190        tensor_1_grad
191            .to_data()
192            .assert_approx_eq::<FT>(&expected, Tolerance::default());
193
194        let expected = TensorData::zeros::<f32, _>(tensor_2_grad.shape());
195        tensor_2_grad
196            .to_data()
197            .assert_approx_eq::<FT>(&expected, Tolerance::default());
198    }
199
200    #[test]
201    fn display() {
202        let config = LayerNormConfig::new(6);
203        let layer_norm = config.init::<TestBackend>(&Default::default());
204
205        assert_eq!(
206            format!("{layer_norm}"),
207            "LayerNorm {d_model: 6, epsilon: 0.00001, params: 12}"
208        );
209    }
210}