rlx_optim/lamb.rs
1// RLX — versatile ML compiler + runtime.
2// Copyright (C) 2026 Eugene Hauptmann, Nataliya Kosmyna.
3// SPDX-License-Identifier: MIT OR Apache-2.0
4
5//! LAMB — Layer-wise Adaptive Moments for Batch training (You et al.,
6//! 2019, "Large Batch Optimization for Deep Learning: Training BERT
7//! in 76 minutes").
8//!
9//! # Idea
10//!
11//! Naïve large-batch training stalls because the per-coordinate Adam
12//! step doesn't account for the magnitude difference between
13//! different layers' weights. LAMB rescales each tensor's Adam-style
14//! update by the **trust ratio** `‖θ‖ / ‖u‖`, so that the per-step
15//! relative change `‖Δθ‖ / ‖θ‖` is bounded and identical across layers.
16//!
17//! # Update rule
18//!
19//! For each tensor (and its flat parameter vector θ):
20//!
21//! ```text
22//! m_t = β₁·m_{t-1} + (1 − β₁)·g_t
23//! v_t = β₂·v_{t-1} + (1 − β₂)·g_t²
24//! u_t = m̂_t / (√v̂_t + ε) + λ·θ_{t-1} // raw update
25//! r_t = ‖θ_{t-1}‖₂ / ‖u_t‖₂ // trust ratio
26//! θ_t = θ_{t-1} − lr · r_t · u_t
27//! ```
28//!
29//! `r_t` is clamped to `1.0` when either norm is zero (warm-up edge
30//! case). LAMB's headline result is that this rescaling makes very
31//! large batch sizes (32k–64k) viable without quality loss.
32//!
33//! # When to use
34//!
35//! Large-batch pre-training — BERT/ViT/Llama-scale data-parallel
36//! runs. State cost = Adam (two buffers).
37
38use std::collections::HashMap;
39
40use crate::Optimizer;
41use crate::common::{l2_norm, zeros_entry};
42
43/// Layer-wise Adaptive Moments for Batch training.
44///
45/// Per-tensor state: two `f32` buffers + a per-call scratch buffer
46/// for the trust-ratio numerator (allocated inside [`Optimizer::step`]).
47#[derive(Debug, Clone)]
48pub struct Lamb {
49 /// Learning rate.
50 pub lr: f32,
51 /// First-moment EMA decay β₁. Default `0.9`.
52 pub beta1: f32,
53 /// Second-moment EMA decay β₂. Default `0.999`.
54 pub beta2: f32,
55 /// Denominator stability constant. Default `1e-6` (looser than
56 /// Adam's `1e-8` — matches NVIDIA's reference).
57 pub eps: f32,
58 /// Decoupled weight-decay coefficient λ. Default `0.01`.
59 pub weight_decay: f32,
60 /// If `true`, divide by bias-corrected moments. Defaults to `true`
61 /// (matches NVIDIA's reference impl); the original paper omits it.
62 pub bias_correction: bool,
63 step: u64,
64 m: HashMap<String, Vec<f32>>,
65 v: HashMap<String, Vec<f32>>,
66 /// Reusable per-tensor scratch buffer for the trust-ratio
67 /// numerator. Cached so we don't allocate every step.
68 scratch: HashMap<String, Vec<f32>>,
69}
70
71impl Lamb {
72 /// Construct with `(β₁, β₂, ε, λ) = (0.9, 0.999, 1e-6, 0.01)`.
73 pub fn new(lr: f32) -> Self {
74 Self {
75 lr,
76 beta1: 0.9,
77 beta2: 0.999,
78 eps: 1e-6,
79 weight_decay: 0.01,
80 bias_correction: true,
81 step: 0,
82 m: HashMap::new(),
83 v: HashMap::new(),
84 scratch: HashMap::new(),
85 }
86 }
87
88 /// Override (β₁, β₂).
89 pub fn with_betas(mut self, b1: f32, b2: f32) -> Self {
90 self.beta1 = b1;
91 self.beta2 = b2;
92 self
93 }
94
95 /// Override the decoupled-decay coefficient.
96 pub fn with_weight_decay(mut self, wd: f32) -> Self {
97 self.weight_decay = wd;
98 self
99 }
100}
101
102impl Optimizer for Lamb {
103 fn set_lr(&mut self, lr: f32) {
104 self.lr = lr;
105 }
106
107 fn step(&mut self, name: &str, _shape: &[usize], param: &mut [f32], grad: &[f32]) {
108 debug_assert_eq!(param.len(), grad.len());
109 let t = (self.step + 1) as f64;
110 let b1 = self.beta1 as f64;
111 let b2 = self.beta2 as f64;
112 let (bc1, bc2) = if self.bias_correction {
113 (1.0 - b1.powf(t), 1.0 - b2.powf(t))
114 } else {
115 (1.0, 1.0)
116 };
117 let eps = self.eps as f64;
118 let lr = self.lr;
119 let wd = self.weight_decay as f64;
120 let m = zeros_entry(&mut self.m, name, param.len());
121 let v = zeros_entry(&mut self.v, name, param.len());
122 let update = zeros_entry(&mut self.scratch, name, param.len());
123 // First pass: update m/v, build `r_i = m_hat / (sqrt(v_hat) + eps) + wd * w`.
124 for i in 0..param.len() {
125 let g = grad[i] as f64;
126 let mi = b1 * m[i] as f64 + (1.0 - b1) * g;
127 let vi = b2 * v[i] as f64 + (1.0 - b2) * g * g;
128 m[i] = mi as f32;
129 v[i] = vi as f32;
130 let m_hat = mi / bc1;
131 let v_hat = vi / bc2;
132 update[i] = (m_hat / (v_hat.sqrt() + eps) + wd * param[i] as f64) as f32;
133 }
134 let w_norm = l2_norm(param);
135 let r_norm = l2_norm(update);
136 let trust = if w_norm > 0.0 && r_norm > 0.0 {
137 w_norm / r_norm
138 } else {
139 1.0
140 };
141 let step_size = lr * trust;
142 for i in 0..param.len() {
143 param[i] -= step_size * update[i];
144 }
145 }
146
147 fn end_iteration(&mut self) {
148 self.step += 1;
149 }
150}