rlx_optim/adam.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//! Adam — Adaptive Moment Estimation (Kingma & Ba, 2014).
6//!
7//! # Update rule
8//!
9//! For each parameter, with `t` the 1-based iteration index:
10//!
11//! ```text
12//! g_t = ∇L(θ_{t-1}) + λ·θ_{t-1} // L2 decay folded in
13//! m_t = β₁·m_{t-1} + (1 − β₁)·g_t
14//! v_t = β₂·v_{t-1} + (1 − β₂)·g_t²
15//! m̂_t = m_t / (1 − β₁ᵗ) // bias correction
16//! v̂_t = v_t / (1 − β₂ᵗ)
17//! θ_t = θ_{t-1} − lr · m̂_t / (√v̂_t + ε)
18//! ```
19//!
20//! # When to use
21//!
22//! Reliable default for transformer pre-training and most non-vision
23//! workloads. Per-parameter memory is **2×** the parameter size (one
24//! `m`, one `v`). If you need decoupled weight decay (recommended for
25//! transformers) use [`crate::AdamW`] instead.
26
27use std::collections::HashMap;
28
29use crate::Optimizer;
30use crate::common::{zeros_entry, zip4_for_each};
31
32/// Bias-corrected first/second moment optimizer.
33///
34/// Per-tensor state: two `f32` buffers (`m`, `v`) of the same shape as
35/// the parameter.
36#[derive(Debug, Clone)]
37pub struct Adam {
38 /// Learning rate. Typical: `1e-3` for from-scratch CNNs, `1e-4`
39 /// for transformer fine-tuning.
40 pub lr: f32,
41 /// First-moment EMA decay β₁ ∈ \[0, 1). Default `0.9`.
42 pub beta1: f32,
43 /// Second-moment EMA decay β₂ ∈ \[0, 1). Default `0.999`.
44 pub beta2: f32,
45 /// Stability constant in the denominator. Default `1e-8`.
46 pub eps: f32,
47 /// L2 weight decay coefficient. **Folded into the gradient**
48 /// (the "classic Adam" rule); use [`crate::AdamW`] for decoupled
49 /// decay. Default `0.0`.
50 pub weight_decay: f32,
51 /// When true, the moment and parameter updates stay in pure `f32`.
52 /// Default `false`, which uses `f64` intermediates: slightly more accurate,
53 /// and measured **2.2x slower** (0.88 vs 0.40 ns/element) because every
54 /// element pays two conversions and an `f64` square root. Mirrors
55 /// [`crate::AdamW::f32_math`], which had the option while this did not.
56 pub f32_math: bool,
57 step: u64,
58 m: HashMap<String, Vec<f32>>,
59 v: HashMap<String, Vec<f32>>,
60}
61
62impl Adam {
63 /// Construct with the given learning rate and the standard
64 /// (β₁, β₂, ε) = (0.9, 0.999, 1e-8) defaults.
65 /// Keep the moment arithmetic in `f32` (see [`Self::f32_math`]).
66 pub fn with_f32_math(mut self, on: bool) -> Self {
67 self.f32_math = on;
68 self
69 }
70
71 pub fn new(lr: f32) -> Self {
72 Self {
73 lr,
74 beta1: 0.9,
75 beta2: 0.999,
76 eps: 1e-8,
77 weight_decay: 0.0,
78 f32_math: false,
79 step: 0,
80 m: HashMap::new(),
81 v: HashMap::new(),
82 }
83 }
84
85 /// Override (β₁, β₂).
86 pub fn with_betas(mut self, b1: f32, b2: f32) -> Self {
87 self.beta1 = b1;
88 self.beta2 = b2;
89 self
90 }
91
92 /// Override the denominator stability constant ε.
93 pub fn with_eps(mut self, eps: f32) -> Self {
94 self.eps = eps;
95 self
96 }
97
98 /// Override the L2 weight-decay coefficient.
99 pub fn with_weight_decay(mut self, wd: f32) -> Self {
100 self.weight_decay = wd;
101 self
102 }
103
104 /// 1-based iteration counter. Starts at 1 (so the first call to
105 /// `step()` sees `t=1`), advances on [`Optimizer::end_iteration`].
106 pub fn current_step(&self) -> u64 {
107 self.step + 1
108 }
109}
110
111// ── checkpointing ───────────────────────────────────────────
112
113impl Adam {
114 /// Named accumulators plus the step counter — see
115 /// [`crate::OptimizerState`].
116 pub(crate) fn snapshot(&self) -> crate::OptimizerState {
117 let mut out = crate::OptimizerState {
118 step: self.step,
119 buffers: Vec::new(),
120 };
121 out.extend_slot("m", &self.m);
122 out.extend_slot("v", &self.v);
123 out
124 }
125
126 pub(crate) fn restore(&mut self, state: &crate::OptimizerState) {
127 self.step = state.step;
128 state.take_slot("m", &mut self.m);
129 state.take_slot("v", &mut self.v);
130 }
131}
132
133impl Optimizer for Adam {
134 fn set_lr(&mut self, lr: f32) {
135 self.lr = lr;
136 }
137
138 fn state_dict(&self) -> Option<crate::OptimizerState> {
139 Some(self.snapshot())
140 }
141
142 fn load_state_dict(&mut self, state: &crate::OptimizerState) -> bool {
143 self.restore(state);
144 true
145 }
146
147 fn step(&mut self, name: &str, _shape: &[usize], param: &mut [f32], grad: &[f32]) {
148 debug_assert_eq!(param.len(), grad.len());
149 if self.f32_math {
150 let t = (self.step + 1) as i32;
151 let b1 = self.beta1;
152 let b2 = self.beta2;
153 let bc1 = 1.0 - b1.powi(t);
154 let bc2 = 1.0 - b2.powi(t);
155 let eps = self.eps;
156 let lr = self.lr;
157 let wd = self.weight_decay;
158 let n = param.len();
159 let m = zeros_entry(&mut self.m, name, n);
160 let v = zeros_entry(&mut self.v, name, n);
161 zip4_for_each(param, m, v, grad, |p, mi, vi, gi| {
162 let g = gi + wd * *p;
163 let new_m = b1 * *mi + (1.0 - b1) * g;
164 let new_v = b2 * *vi + (1.0 - b2) * (g * g);
165 *mi = new_m;
166 *vi = new_v;
167 let m_hat = new_m / bc1;
168 let v_hat = new_v / bc2;
169 *p -= lr * (m_hat / (v_hat.sqrt() + eps));
170 });
171 return;
172 }
173 let t = (self.step + 1) as f64;
174 let b1 = self.beta1 as f64;
175 let b2 = self.beta2 as f64;
176 let bc1 = 1.0 - b1.powf(t);
177 let bc2 = 1.0 - b2.powf(t);
178 let eps = self.eps as f64;
179 let lr = self.lr as f64;
180 let wd = self.weight_decay;
181 let n = param.len();
182 // `self.m` / `self.v` are distinct fields, so the two
183 // `zeros_entry` calls borrow disjoint regions of `self` and
184 // their results can coexist.
185 let m = zeros_entry(&mut self.m, name, n);
186 let v = zeros_entry(&mut self.v, name, n);
187 zip4_for_each(param, m, v, grad, |p, mi, vi, gi| {
188 let g = (gi + wd * *p) as f64;
189 let new_m = b1 * *mi as f64 + (1.0 - b1) * g;
190 let new_v = b2 * *vi as f64 + (1.0 - b2) * g * g;
191 *mi = new_m as f32;
192 *vi = new_v as f32;
193 let m_hat = new_m / bc1;
194 let v_hat = new_v / bc2;
195 *p -= (lr * m_hat / (v_hat.sqrt() + eps)) as f32;
196 });
197 }
198
199 fn end_iteration(&mut self) {
200 self.step += 1;
201 }
202}