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#[derive(Config, Debug)]
17pub struct LambConfig {
18 #[config(default = 0.9)]
20 beta_1: f32,
21 #[config(default = 0.999)]
23 beta_2: f32,
24 #[config(default = 1e-6)]
26 epsilon: f32,
27 #[config(default = 0.0)]
29 weight_decay: f32,
30 #[config(default = true)]
32 use_trust_ratio: bool,
33 grad_clipping: Option<GradientClippingConfig>,
35}
36
37#[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#[derive(RecordState, Clone)]
56pub struct LambState<const D: usize> {
57 pub time: usize,
59 pub moment_1: Tensor<D>,
61 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 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 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 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}