Skip to main content

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}