Skip to main content

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}