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}