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}