Skip to main content

rlx_optim/
adamw.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//! AdamW — Adam with decoupled weight decay (Loshchilov & Hutter, 2017).
6//!
7//! # Why "decoupled"?
8//!
9//! In classical Adam, an L2 penalty `λ·θ` is folded into the gradient
10//! before the EMAs — but the second-moment estimator `v_t` then scales
11//! the decay term, so parameters with large gradients get *less* L2
12//! regularization. AdamW separates the two: decay multiplies the
13//! parameter directly, *outside* the adaptive step.
14//!
15//! # Update rule
16//!
17//! ```text
18//! m_t = β₁·m_{t-1} + (1 − β₁)·g_t            // g_t is the raw grad
19//! v_t = β₂·v_{t-1} + (1 − β₂)·g_t²
20//! m̂_t = m_t / (1 − β₁ᵗ)
21//! v̂_t = v_t / (1 − β₂ᵗ)
22//! θ_t = θ_{t-1} − lr · ( m̂_t/(√v̂_t + ε) + λ·θ_{t-1} )
23//! ```
24//!
25//! # When to use
26//!
27//! The default for transformer pre-training. Pair with a cosine /
28//! linear LR schedule and `weight_decay = 0.1` for LLMs, `0.01–0.05`
29//! for vision transformers.
30
31use std::collections::HashMap;
32
33use crate::Optimizer;
34use crate::common::{zeros_entry, zip4_for_each};
35
36/// Adam with decoupled weight decay.
37///
38/// Per-tensor state identical to [`crate::Adam`] (two `f32` buffers).
39#[derive(Debug, Clone)]
40pub struct AdamW {
41    /// Learning rate. Typical LLM pre-training value: `1e-4` to `3e-4`.
42    pub lr: f32,
43    /// First-moment EMA decay. Default `0.9`.
44    pub beta1: f32,
45    /// Second-moment EMA decay. Default `0.999` (matches Adam);
46    /// `0.95` is common for very long pre-training runs.
47    pub beta2: f32,
48    /// Denominator stability constant. Default `1e-8`.
49    pub eps: f32,
50    /// **Decoupled** weight-decay coefficient λ. Multiplies the
51    /// parameter directly inside the update; `0.01–0.1` typical.
52    /// Defaults to `0.01`.
53    pub weight_decay: f32,
54    /// When true, moment/parameter updates stay in pure `f32` (PyTorch AdamW /
55    /// `foreach=False` default). Default `false` uses `f64` intermediates
56    /// (slightly more accurate, but drifts vs PyTorch shared-init checks).
57    pub f32_math: bool,
58    step: u64,
59    m: HashMap<String, Vec<f32>>,
60    v: HashMap<String, Vec<f32>>,
61}
62
63impl AdamW {
64    /// Construct with the given learning rate and the standard
65    /// (β₁, β₂, ε, λ) = (0.9, 0.999, 1e-8, 0.01) defaults.
66    pub fn new(lr: f32) -> Self {
67        Self {
68            lr,
69            beta1: 0.9,
70            beta2: 0.999,
71            eps: 1e-8,
72            weight_decay: 0.01,
73            f32_math: false,
74            step: 0,
75            m: HashMap::new(),
76            v: HashMap::new(),
77        }
78    }
79
80    /// Override (β₁, β₂).
81    pub fn with_betas(mut self, b1: f32, b2: f32) -> Self {
82        self.beta1 = b1;
83        self.beta2 = b2;
84        self
85    }
86
87    /// Override the decoupled-decay coefficient.
88    pub fn with_weight_decay(mut self, wd: f32) -> Self {
89        self.weight_decay = wd;
90        self
91    }
92
93    /// Override the denominator ε.
94    pub fn with_eps(mut self, eps: f32) -> Self {
95        self.eps = eps;
96        self
97    }
98
99    /// Pure-`f32` AdamW arithmetic (match PyTorch training trajectories).
100    pub fn with_f32_math(mut self, on: bool) -> Self {
101        self.f32_math = on;
102        self
103    }
104}
105
106// ── checkpointing ───────────────────────────────────────────
107
108impl AdamW {
109    /// Named accumulators plus the step counter — see
110    /// [`crate::OptimizerState`].
111    pub(crate) fn snapshot(&self) -> crate::OptimizerState {
112        let mut out = crate::OptimizerState {
113            step: self.step,
114            buffers: Vec::new(),
115        };
116        out.extend_slot("m", &self.m);
117        out.extend_slot("v", &self.v);
118        out
119    }
120
121    pub(crate) fn restore(&mut self, state: &crate::OptimizerState) {
122        self.step = state.step;
123        state.take_slot("m", &mut self.m);
124        state.take_slot("v", &mut self.v);
125    }
126}
127
128impl Optimizer for AdamW {
129    fn set_lr(&mut self, lr: f32) {
130        self.lr = lr;
131    }
132
133    fn state_dict(&self) -> Option<crate::OptimizerState> {
134        Some(self.snapshot())
135    }
136
137    fn load_state_dict(&mut self, state: &crate::OptimizerState) -> bool {
138        self.restore(state);
139        true
140    }
141
142    fn step(&mut self, name: &str, _shape: &[usize], param: &mut [f32], grad: &[f32]) {
143        debug_assert_eq!(param.len(), grad.len());
144        let m = zeros_entry(&mut self.m, name, param.len());
145        let v = zeros_entry(&mut self.v, name, param.len());
146        if self.f32_math {
147            let t = (self.step + 1) as i32;
148            let b1 = self.beta1;
149            let b2 = self.beta2;
150            let bc1 = 1.0 - b1.powi(t);
151            let bc2 = 1.0 - b2.powi(t);
152            let eps = self.eps;
153            let lr = self.lr;
154            let wd = self.weight_decay;
155            zip4_for_each(param, m, v, grad, |p, mi, vi, gi| {
156                let new_m = b1 * *mi + (1.0 - b1) * gi;
157                let new_v = b2 * *vi + (1.0 - b2) * (gi * gi);
158                *mi = new_m;
159                *vi = new_v;
160                let m_hat = new_m / bc1;
161                let v_hat = new_v / bc2;
162                // PyTorch AdamW: param *= (1 - lr*wd); param += -lr * m_hat/(sqrt(v_hat)+eps)
163                *p = *p * (1.0 - lr * wd) - lr * (m_hat / (v_hat.sqrt() + eps));
164            });
165        } else {
166            let t = (self.step + 1) as f64;
167            let b1 = self.beta1 as f64;
168            let b2 = self.beta2 as f64;
169            let bc1 = 1.0 - b1.powf(t);
170            let bc2 = 1.0 - b2.powf(t);
171            let eps = self.eps as f64;
172            let lr = self.lr as f64;
173            let wd = self.weight_decay as f64;
174            zip4_for_each(param, m, v, grad, |p, mi, vi, gi| {
175                let g = gi as f64;
176                let new_m = b1 * *mi as f64 + (1.0 - b1) * g;
177                let new_v = b2 * *vi as f64 + (1.0 - b2) * g * g;
178                *mi = new_m as f32;
179                *vi = new_v as f32;
180                let m_hat = new_m / bc1;
181                let v_hat = new_v / bc2;
182                // Decoupled decay: applied to the parameter, then the
183                // adaptive step is subtracted.
184                let pf = *p as f64;
185                *p = (pf - lr * (m_hat / (v_hat.sqrt() + eps) + wd * pf)) as f32;
186            });
187        }
188    }
189
190    fn end_iteration(&mut self) {
191        self.step += 1;
192    }
193}