Skip to main content

rlx_optim/
adafactor.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//! Adafactor (Shazeer & Stern, 2018, "Adafactor: Adaptive Learning
6//! Rates with Sublinear Memory Cost").
7//!
8//! # Idea
9//!
10//! Adam's `v_t` is the same shape as θ — for a 70B-parameter model
11//! that's 280 GB of optimizer state. Adafactor *factorizes* the
12//! second-moment matrix: for a 2-D parameter of shape `m × n`, instead
13//! of an `m·n` buffer it stores a row-statistic `R ∈ ℝᵐ` and a
14//! column-statistic `C ∈ ℝⁿ`, then reconstructs
15//! `V̂_{ij} ≈ R_i · C_j / Σ_k R_k`. State drops from `O(m·n)` to
16//! `O(m + n)`.
17//!
18//! # Update rule (this impl: factored 2nd-moment, no 1st-moment)
19//!
20//! Let `β₂_t = 1 − t^{decay_rate}` (default decay_rate = −0.8). For a
21//! 2-D parameter:
22//!
23//! ```text
24//! R_i = β₂_t·R_i + (1−β₂_t)·mean_j(g_ij² + ε₁)
25//! C_j = β₂_t·C_j + (1−β₂_t)·mean_i(g_ij² + ε₁)
26//! V̂_{ij} = R_i · C_j / Σ_k R_k
27//! u_{ij}  = g_{ij} / √V̂_{ij}
28//! u ← u / max(1, RMS(u) / clip_threshold)        // RMS-of-update clip
29//! lr_t    = manual_lr OR  min(1/√t, 1e-2) · max(ε₂, RMS(θ))   // relative step
30//! θ_t     = θ_{t-1} − lr_t · ( u + λ·θ_{t-1} )
31//! ```
32//!
33//! For non-2-D parameters (bias vectors, 4-D conv weights) we fall
34//! back to a full per-element EMA — the savings are negligible there
35//! anyway. The optional first-moment EMA is **not** implemented
36//! (matches the recommended T5 configuration).
37//!
38//! # When to use
39//!
40//! When you don't have memory for Adam-style optimizer state — large
41//! models, low-VRAM fine-tuning, sequence-length scaling experiments.
42//! State cost per matrix = `m + n` floats vs Adam's `2·m·n`.
43
44use std::collections::HashMap;
45
46use crate::Optimizer;
47use crate::common::{l2_norm, zeros_entry};
48
49/// Adafactor — factored-second-moment optimizer.
50///
51/// Per-tensor state: a `rows`-vector + a `cols`-vector for 2-D
52/// parameters (sublinear in `rows·cols`), or a full EMA for non-2-D.
53#[derive(Debug, Clone)]
54pub struct Adafactor {
55    /// Optional manual learning rate. `None` ⇒ use the "relative
56    /// step" rule `min(1/√t, 1e-2) · max(ε₂, RMS(θ))` from the paper.
57    /// Default `None`.
58    pub lr: Option<f32>,
59    /// β₂_t decay-rate exponent. `β₂_t = 1 − tˣ` with `x = -0.8`
60    /// (default) means slow decay early, full decay asymptotically.
61    pub beta2_decay: f32,
62    /// Squared-gradient stability constant added before each row /
63    /// column average. Default `1e-30`.
64    pub eps1: f32,
65    /// RMS-of-parameter floor for the relative-step rule. Default `1e-3`.
66    pub eps2: f32,
67    /// Update-RMS clipping threshold (Shazeer & Stern §6). Default `1.0`.
68    pub clip_threshold: f32,
69    /// Decoupled weight-decay coefficient λ. Default `0.0`.
70    pub weight_decay: f32,
71    step: u64,
72    // Per-parameter state.
73    r: HashMap<String, Vec<f32>>, // row factor (length rows) for 2D
74    c: HashMap<String, Vec<f32>>, // col factor (length cols) for 2D
75    v: HashMap<String, Vec<f32>>, // full EMA for non-2D
76}
77
78impl Adafactor {
79    /// Construct with paper defaults (no manual lr ⇒ relative step,
80    /// `decay_rate = -0.8`, `ε₁=1e-30, ε₂=1e-3, clip=1.0, λ=0.0`).
81    pub fn new() -> Self {
82        Self {
83            lr: None,
84            beta2_decay: -0.8,
85            eps1: 1e-30,
86            eps2: 1e-3,
87            clip_threshold: 1.0,
88            weight_decay: 0.0,
89            step: 0,
90            r: HashMap::new(),
91            c: HashMap::new(),
92            v: HashMap::new(),
93        }
94    }
95
96    /// Switch from the relative-step rule to a manual learning rate.
97    pub fn with_lr(mut self, lr: f32) -> Self {
98        self.lr = Some(lr);
99        self
100    }
101
102    /// Override the decoupled-decay coefficient.
103    pub fn with_weight_decay(mut self, wd: f32) -> Self {
104        self.weight_decay = wd;
105        self
106    }
107}
108
109impl Default for Adafactor {
110    fn default() -> Self {
111        Self::new()
112    }
113}
114
115impl Optimizer for Adafactor {
116    fn step(&mut self, name: &str, shape: &[usize], param: &mut [f32], grad: &[f32]) {
117        debug_assert_eq!(param.len(), grad.len());
118        let t = (self.step + 1) as f64;
119        // β₂_t = 1 − t^{beta2_decay}, decay_rate ∈ (-1, 0).
120        let beta2_t = 1.0 - t.powf(self.beta2_decay as f64);
121        let eps1 = self.eps1 as f64;
122        let clip = self.clip_threshold as f64;
123        let n = param.len();
124
125        // ── Update second-moment estimate ──────────────────────────
126        let mut update = vec![0.0f32; n];
127        if shape.len() == 2 {
128            let (rows, cols) = (shape[0], shape[1]);
129            debug_assert_eq!(rows * cols, n);
130            let r = zeros_entry(&mut self.r, name, rows);
131            // Row factor: average of g² across columns, then EMA.
132            let mut row_buf = vec![0.0f64; rows];
133            for i in 0..rows {
134                let mut s = 0.0f64;
135                for j in 0..cols {
136                    let g = grad[i * cols + j] as f64;
137                    s += g * g + eps1;
138                }
139                row_buf[i] = s / cols as f64;
140            }
141            for i in 0..rows {
142                r[i] = (beta2_t * r[i] as f64 + (1.0 - beta2_t) * row_buf[i]) as f32;
143            }
144            let r_snapshot: Vec<f32> = r.clone();
145
146            // Column factor: average of g² across rows, then EMA.
147            let c = zeros_entry(&mut self.c, name, cols);
148            let mut col_buf = vec![0.0f64; cols];
149            for j in 0..cols {
150                let mut s = 0.0f64;
151                for i in 0..rows {
152                    let g = grad[i * cols + j] as f64;
153                    s += g * g + eps1;
154                }
155                col_buf[j] = s / rows as f64;
156            }
157            for j in 0..cols {
158                c[j] = (beta2_t * c[j] as f64 + (1.0 - beta2_t) * col_buf[j]) as f32;
159            }
160            let r_sum: f64 = r_snapshot.iter().map(|&x| x as f64).sum();
161            // v_ij = r_i * c_j / (sum_k r_k). Build update = g / sqrt(v).
162            for i in 0..rows {
163                for j in 0..cols {
164                    let v_ij = r_snapshot[i] as f64 * c[j] as f64 / r_sum.max(eps1);
165                    let g = grad[i * cols + j] as f64;
166                    update[i * cols + j] = (g / v_ij.sqrt().max(eps1.sqrt())) as f32;
167                }
168            }
169        } else {
170            // Non-2D: full per-element EMA.
171            let v = zeros_entry(&mut self.v, name, n);
172            for i in 0..n {
173                let g = grad[i] as f64;
174                v[i] = (beta2_t * v[i] as f64 + (1.0 - beta2_t) * (g * g + eps1)) as f32;
175                update[i] = (g / (v[i] as f64).sqrt().max(eps1.sqrt())) as f32;
176            }
177        }
178
179        // RMS-of-update clipping (Shazeer & Stern §6).
180        let u_rms = (l2_norm(&update) as f64 / (n as f64).sqrt()).max(1.0 / clip);
181        let scale = (1.0 / (u_rms * clip)).min(1.0);
182        for u in update.iter_mut() {
183            *u = (*u as f64 * scale) as f32;
184        }
185
186        // Learning rate (relative-step or manual).
187        let lr = match self.lr {
188            Some(x) => x as f64,
189            None => {
190                let p_rms = (l2_norm(param) as f64 / (n as f64).sqrt()).max(self.eps2 as f64);
191                (1.0 / t.sqrt()).min(1e-2) * p_rms
192            }
193        };
194        let wd = self.weight_decay as f64;
195        for i in 0..n {
196            let p = param[i] as f64;
197            param[i] = (p - lr * (update[i] as f64 + wd * p)) as f32;
198        }
199    }
200
201    fn end_iteration(&mut self) {
202        self.step += 1;
203    }
204}