Skip to main content

burn_optim/optim/
lamb.rs

1use burn_core as burn;
2
3use crate::{LearningRate, RecordState, grad_clipping::GradientClippingConfig};
4use burn::{
5    config::Config,
6    tensor::{Device, FloatDType, Tensor},
7};
8
9use super::{Optimizer, module_optimizer::ModuleOptimizer};
10
11#[cfg(not(feature = "std"))]
12#[allow(unused_imports)]
13use num_traits::Float as _;
14
15/// [`Lamb`] configuration.
16#[derive(Config, Debug)]
17pub struct LambConfig {
18    /// Exponential decay rate for the first moment estimates.
19    #[config(default = 0.9)]
20    beta_1: f32,
21    /// Exponential decay rate for the second moment estimates.
22    #[config(default = 0.999)]
23    beta_2: f32,
24    /// A value added to the denominator for numerical stability.
25    #[config(default = 1e-6)]
26    epsilon: f32,
27    /// Weight decay applied to the Adam update before layer-wise adaptation.
28    #[config(default = 0.0)]
29    weight_decay: f32,
30    /// Whether to scale each parameter update by its layer-wise trust ratio.
31    #[config(default = true)]
32    use_trust_ratio: bool,
33    /// Optional gradient clipping configuration.
34    grad_clipping: Option<GradientClippingConfig>,
35}
36
37/// Layer-wise Adaptive Moments (LAMB) optimizer.
38///
39/// LAMB applies an Adam-style update, then scales it by the ratio between the parameter norm and
40/// update norm. This layer-wise adaptation is useful when training with very large batch sizes.
41///
42/// See: [Large Batch Optimization for Deep Learning: Training BERT in 76 minutes](https://arxiv.org/abs/1904.00962).
43///
44/// Configured by [`LambConfig`].
45#[derive(Clone)]
46pub struct Lamb {
47    beta_1: f32,
48    beta_2: f32,
49    epsilon: f32,
50    weight_decay: f32,
51    use_trust_ratio: bool,
52}
53
54/// LAMB state for a single parameter tensor.
55#[derive(RecordState, Clone)]
56pub struct LambState<const D: usize> {
57    /// The number of optimization steps applied to the parameter.
58    pub time: usize,
59    /// Exponential moving average of gradients.
60    pub moment_1: Tensor<D>,
61    /// Exponential moving average of squared gradients.
62    pub moment_2: Tensor<D>,
63}
64
65impl Optimizer for Lamb {
66    type State<const D: usize> = LambState<D>;
67
68    fn step<const D: usize>(
69        &self,
70        lr: LearningRate,
71        tensor: Tensor<D>,
72        grad: Tensor<D>,
73        state: Option<Self::State<D>>,
74    ) -> (Tensor<D>, Option<Self::State<D>>) {
75        let factor_1 = 1.0 - self.beta_1;
76        let factor_2 = 1.0 - self.beta_2;
77
78        let state = if let Some(mut state) = state {
79            state.moment_1 = state
80                .moment_1
81                .mul_scalar(self.beta_1)
82                .add(grad.clone().mul_scalar(factor_1));
83            state.moment_2 = state
84                .moment_2
85                .mul_scalar(self.beta_2)
86                .add(grad.square().mul_scalar(factor_2));
87            state.time += 1;
88            state
89        } else {
90            LambState {
91                time: 1,
92                moment_1: grad.clone().mul_scalar(factor_1),
93                moment_2: grad.square().mul_scalar(factor_2),
94            }
95        };
96
97        let time = state.time as i32;
98        let moment_1 = state
99            .moment_1
100            .clone()
101            .div_scalar(1.0 - self.beta_1.powi(time));
102        let moment_2 = state
103            .moment_2
104            .clone()
105            .div_scalar(1.0 - self.beta_2.powi(time));
106
107        let mut update = moment_1.div(moment_2.sqrt().add_scalar(self.epsilon));
108        if self.weight_decay != 0.0 {
109            update = update.add(tensor.clone().mul_scalar(self.weight_decay));
110        }
111
112        let update = if self.use_trust_ratio {
113            let parameter_norm = l2_norm(tensor.clone());
114            let update_norm = l2_norm(update.clone());
115            let valid_norms = parameter_norm
116                .clone()
117                .greater_scalar(0.0)
118                .bool_and(update_norm.clone().greater_scalar(0.0));
119
120            // Avoid forming 0 / 0 even though the invalid result would be masked out below.
121            let min_positive = update
122                .dtype()
123                .finfo()
124                .unwrap_or(FloatDType::F32.finfo())
125                .min_positive;
126            let ratio = parameter_norm.div(update_norm.clamp_min(min_positive));
127            let trust_ratio = ratio.ones_like().mask_where(valid_norms, ratio);
128
129            update.mul(trust_ratio.unsqueeze())
130        } else {
131            update
132        };
133
134        let tensor = tensor - update.mul_scalar(lr);
135        (tensor, Some(state))
136    }
137
138    fn to_device<const D: usize>(mut state: Self::State<D>, device: &Device) -> Self::State<D> {
139        state.moment_1 = state.moment_1.to_device(device);
140        state.moment_2 = state.moment_2.to_device(device);
141        state
142    }
143}
144
145impl LambConfig {
146    /// Build the per-parameter LAMB optimizer.
147    ///
148    /// Use [`Self::init`] to construct a whole-module optimizer with the configured gradient
149    /// clipping behavior.
150    pub fn build(&self) -> Lamb {
151        Lamb {
152            beta_1: self.beta_1,
153            beta_2: self.beta_2,
154            epsilon: self.epsilon,
155            weight_decay: self.weight_decay,
156            use_trust_ratio: self.use_trust_ratio,
157        }
158    }
159
160    /// Initialize a whole-module LAMB optimizer.
161    pub fn init(&self) -> ModuleOptimizer {
162        let mut optimizer = ModuleOptimizer::from(self.build());
163        if let Some(config) = &self.grad_clipping {
164            optimizer = optimizer.with_grad_clipping(config.init());
165        }
166        optimizer
167    }
168}
169
170fn l2_norm<const D: usize>(tensor: Tensor<D>) -> Tensor<1> {
171    tensor.square().sum().sqrt()
172}
173
174#[cfg(test)]
175mod tests {
176    use super::*;
177    use crate::GradientsParams;
178    use burn::{
179        module::Param,
180        tensor::{TensorData, Tolerance},
181    };
182    use burn_nn::Linear;
183
184    #[test]
185    fn test_lamb_matches_pytorch_reference_for_two_steps() {
186        let device = Device::default();
187        let optimizer = LambConfig::new()
188            .with_beta_1(0.9)
189            .with_beta_2(0.999)
190            .with_epsilon(1e-6)
191            .with_weight_decay(0.1)
192            .build();
193        let tensor = Tensor::<1>::from_floats([1.0, -2.0, 3.0], &device);
194
195        let (tensor, state) = optimizer.step(
196            0.01,
197            tensor,
198            Tensor::from_floats([0.1, -0.2, 0.3], &device),
199            None,
200        );
201        tensor.to_data().assert_approx_eq::<f32>(
202            &TensorData::from([0.9802435, -1.9784473, 2.9766512]),
203            Tolerance::absolute(1e-6),
204        );
205
206        let (tensor, state) = optimizer.step(
207            0.01,
208            tensor,
209            Tensor::from_floats([-0.4, 0.5, -0.6], &device),
210            state,
211        );
212        tensor.to_data().assert_approx_eq::<f32>(
213            &TensorData::from([1.0127187, -1.9956441, 2.9814672]),
214            Tolerance::absolute(1e-6),
215        );
216
217        let state = state.unwrap();
218        assert_eq!(state.time, 2);
219        state.moment_1.to_data().assert_approx_eq::<f32>(
220            &TensorData::from([-0.031, 0.032, -0.033]),
221            Tolerance::absolute(1e-7),
222        );
223        state.moment_2.to_data().assert_approx_eq::<f32>(
224            &TensorData::from([0.00016999, 0.00028996, 0.00044991]),
225            Tolerance::absolute(1e-8),
226        );
227    }
228
229    #[test]
230    fn test_lamb_zero_norm_uses_unit_trust_ratio() {
231        let device = Device::default();
232        let optimizer = LambConfig::new().with_weight_decay(0.1).build();
233        let tensor = Tensor::<1>::zeros([2], &device);
234        let grad = Tensor::<1>::zeros([2], &device);
235
236        let (tensor, _) = optimizer.step(0.01, tensor, grad, None);
237
238        tensor
239            .to_data()
240            .assert_eq(&TensorData::from([0.0f32, 0.0]), true);
241    }
242
243    #[test]
244    fn test_lamb_can_disable_trust_ratio() {
245        let device = Device::default();
246        let optimizer = LambConfig::new()
247            .with_epsilon(1e-6)
248            .with_weight_decay(0.1)
249            .with_use_trust_ratio(false)
250            .build();
251        let tensor = Tensor::<1>::from_floats([1.0, -2.0, 3.0], &device);
252        let grad = Tensor::<1>::from_floats([0.1, -0.2, 0.3], &device);
253
254        let (tensor, _) = optimizer.step(0.01, tensor, grad, None);
255
256        tensor.to_data().assert_approx_eq::<f32>(
257            &TensorData::from([0.9890001, -1.988, 2.987]),
258            Tolerance::absolute(1e-6),
259        );
260    }
261
262    #[test]
263    fn test_lamb_state_survives_burnpack_round_trip() {
264        let device = Device::default().autodiff();
265        let linear = Linear {
266            weight: Param::from_data(TensorData::from([[1.0, -2.0], [3.0, -4.0]]), &device),
267            bias: Some(Param::from_data(TensorData::from([0.5, -0.5]), &device)),
268        };
269        let input = Tensor::<2>::from_floats([[0.25, -0.75]], &device).require_grad();
270        let mut optimizer = LambConfig::new().with_weight_decay(0.1).init();
271
272        let grads = GradientsParams::from_grads(linear.forward(input.clone()).backward(), &linear);
273        let linear = optimizer.step(0.01, linear, grads);
274        let bytes = optimizer.into_bytes().unwrap();
275        assert!(!bytes.is_empty());
276
277        let mut reloaded = LambConfig::new()
278            .with_weight_decay(0.1)
279            .init()
280            .from_bytes(bytes)
281            .unwrap();
282        let grads_original =
283            GradientsParams::from_grads(linear.forward(input.clone()).backward(), &linear);
284        let grads_reloaded = GradientsParams::from_grads(linear.forward(input).backward(), &linear);
285
286        let from_original = optimizer.step(0.01, linear.clone(), grads_original);
287        let from_reloaded = reloaded.step(0.01, linear, grads_reloaded);
288
289        from_original
290            .weight
291            .to_data()
292            .assert_approx_eq::<f32>(&from_reloaded.weight.to_data(), Tolerance::absolute(1e-6));
293        from_original
294            .bias
295            .unwrap()
296            .to_data()
297            .assert_approx_eq::<f32>(
298                &from_reloaded.bias.unwrap().to_data(),
299                Tolerance::absolute(1e-6),
300            );
301    }
302}