rlx_optim/soap.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//! SOAP — ShampoO with Adam in the Preconditioner's eigenbasis
6//! (Vyas, Morwani, Anil, et al., 2024).
7//!
8//! # Idea
9//!
10//! Shampoo (Gupta et al. 2018) preconditions a 2-D parameter's
11//! gradient by `L⁻¹ᐟ⁴ · G · R⁻¹ᐟ⁴`, where `L = E[G·Gᵀ]` and
12//! `R = E[Gᵀ·G]` are Kronecker-factor covariances. SOAP observes that
13//! the same preconditioner is equivalent to **running Adam in the
14//! eigenbasis** of `L` and `R` — and that you only need to recompute
15//! the eigenbasis every K steps. This delivers Shampoo's quality with
16//! Adam's per-step cost (between recompiles).
17//!
18//! # Update rule (for a 2-D parameter `W ∈ ℝ^{m×n}`)
19//!
20//! ```text
21//! L_t = sb·L_{t-1} + (1−sb)·G·Gᵀ // m×m
22//! R_t = sb·R_{t-1} + (1−sb)·Gᵀ·G // n×n
23//! every K steps: // K = precond_freq
24//! Q_L, _ = eigh(L_t) // m×m eigenbasis
25//! Q_R, _ = eigh(R_t) // n×n eigenbasis
26//! G' = Q_Lᵀ · G · Q_R // rotated gradient
27//! [per-element Adam on G' → U']
28//! U = Q_L · U' · Q_Rᵀ // rotate back
29//! θ_t = θ_{t-1} − lr · ( U + λ·θ_{t-1} )
30//! ```
31//!
32//! For non-2-D parameters we fall back to plain AdamW.
33//!
34//! # When to use
35//!
36//! When you want Shampoo's quality on transformers / large dense
37//! models and can afford the eigendecomposition cost amortized over
38//! `precond_freq` steps. State cost per matrix:
39//! `L (m²) + R (n²) + Q_L (m²) + Q_R (n²) + m_rot (m·n) + v_rot (m·n)`.
40
41use std::collections::HashMap;
42
43use crate::Optimizer;
44use crate::common::{jacobi_eigh_sym, matmul, zeros_entry};
45
46#[derive(Debug, Clone)]
47struct SoapState {
48 l: Vec<f32>, // m × m left covariance
49 r: Vec<f32>, // n × n right covariance
50 ql: Vec<f32>, // m × m eigenbasis (row-major: row i = eigvec i)
51 qr: Vec<f32>, // n × n eigenbasis
52 m_rot: Vec<f32>, // first moment in rotated basis (m·n)
53 v_rot: Vec<f32>, // second moment in rotated basis (m·n)
54 initialized_basis: bool,
55}
56
57/// SOAP — Shampoo-in-Adam-basis optimizer.
58#[derive(Debug, Clone)]
59pub struct Soap {
60 /// Learning rate.
61 pub lr: f32,
62 /// First-moment EMA decay (in the *rotated* basis). Default `0.95`.
63 pub beta1: f32,
64 /// Second-moment EMA decay (in the *rotated* basis). Default `0.95`.
65 pub beta2: f32,
66 /// Decay for the L/R covariance EMAs. Often equal to β₂.
67 pub shampoo_beta: f32,
68 /// Denominator stability constant. Default `1e-8`.
69 pub eps: f32,
70 /// Decoupled weight-decay coefficient λ. Default `0.01`.
71 pub weight_decay: f32,
72 /// Recompute the eigenbasis every `precond_freq` steps. Larger
73 /// values amortize the Jacobi cost but lag the preconditioner.
74 /// Default `10`.
75 pub precond_freq: u64,
76 /// Max Jacobi sweeps per rediagonalization. Default `30`.
77 pub jacobi_sweeps: u32,
78 step: u64,
79 state: HashMap<String, SoapState>,
80 // Fallback Adam state for non-2D parameters.
81 fb_m: HashMap<String, Vec<f32>>,
82 fb_v: HashMap<String, Vec<f32>>,
83}
84
85impl Soap {
86 /// Construct with `(β₁, β₂, sb, ε, λ, freq, sweeps) =
87 /// (0.95, 0.95, 0.95, 1e-8, 0.01, 10, 30)`.
88 pub fn new(lr: f32) -> Self {
89 Self {
90 lr,
91 beta1: 0.95,
92 beta2: 0.95,
93 shampoo_beta: 0.95,
94 eps: 1e-8,
95 weight_decay: 0.01,
96 precond_freq: 10,
97 jacobi_sweeps: 30,
98 step: 0,
99 state: HashMap::new(),
100 fb_m: HashMap::new(),
101 fb_v: HashMap::new(),
102 }
103 }
104
105 /// Override the decoupled-decay coefficient.
106 pub fn with_weight_decay(mut self, wd: f32) -> Self {
107 self.weight_decay = wd;
108 self
109 }
110}
111
112impl Optimizer for Soap {
113 fn set_lr(&mut self, lr: f32) {
114 self.lr = lr;
115 }
116
117 fn step(&mut self, name: &str, shape: &[usize], param: &mut [f32], grad: &[f32]) {
118 debug_assert_eq!(param.len(), grad.len());
119 if shape.len() != 2 {
120 // Fallback: plain AdamW for non-matrix parameters.
121 adamw_fallback(self, name, param, grad);
122 return;
123 }
124 let (m, n) = (shape[0], shape[1]);
125 debug_assert_eq!(m * n, param.len());
126 let t = (self.step + 1) as f64;
127 let b1 = self.beta1 as f64;
128 let b2 = self.beta2 as f64;
129 let bc1 = 1.0 - b1.powf(t);
130 let bc2 = 1.0 - b2.powf(t);
131 let sb = self.shampoo_beta as f64;
132 let eps = self.eps;
133 let lr = self.lr;
134 let wd = self.weight_decay;
135
136 let st = self
137 .state
138 .entry(name.to_owned())
139 .or_insert_with(|| SoapState {
140 l: vec![0.0; m * m],
141 r: vec![0.0; n * n],
142 ql: identity(m),
143 qr: identity(n),
144 m_rot: vec![0.0; m * n],
145 v_rot: vec![0.0; m * n],
146 initialized_basis: false,
147 });
148
149 // ── 1. Update L, R covariances ─────────────────────────────
150 // L += (1-sb)·G·Gᵀ; R += (1-sb)·Gᵀ·G (with β2 decay).
151 for i in 0..m {
152 for j in 0..m {
153 let mut s = 0.0f64;
154 for p in 0..n {
155 s += grad[i * n + p] as f64 * grad[j * n + p] as f64;
156 }
157 let lij = sb * st.l[i * m + j] as f64 + (1.0 - sb) * s;
158 st.l[i * m + j] = lij as f32;
159 }
160 }
161 for i in 0..n {
162 for j in 0..n {
163 let mut s = 0.0f64;
164 for p in 0..m {
165 s += grad[p * n + i] as f64 * grad[p * n + j] as f64;
166 }
167 let rij = sb * st.r[i * n + j] as f64 + (1.0 - sb) * s;
168 st.r[i * n + j] = rij as f32;
169 }
170 }
171
172 // ── 2. Rediagonalize periodically (and once on the first step). ──
173 let need_rediag = !st.initialized_basis || self.step.is_multiple_of(self.precond_freq);
174 if need_rediag {
175 let mut l_copy = st.l.clone();
176 let mut r_copy = st.r.clone();
177 jacobi_eigh_sym(&mut l_copy, m, &mut st.ql, self.jacobi_sweeps, 1e-6);
178 jacobi_eigh_sym(&mut r_copy, n, &mut st.qr, self.jacobi_sweeps, 1e-6);
179 st.initialized_basis = true;
180 }
181
182 // ── 3. Rotate gradient: G' = Qₗᵀ · G · Q_r ────────────────
183 let mut tmp = vec![0.0f32; m * n];
184 // tmp = Qₗᵀ · G ⇒ tmp[i,j] = sum_p Qₗ[p,i] · G[p,j]
185 for i in 0..m {
186 for j in 0..n {
187 let mut s = 0.0f32;
188 for p in 0..m {
189 s += st.ql[p * m + i] * grad[p * n + j];
190 }
191 tmp[i * n + j] = s;
192 }
193 }
194 let mut g_rot = vec![0.0f32; m * n];
195 matmul(&tmp, &st.qr, m, n, n, &mut g_rot);
196
197 // ── 4. Per-element Adam on rotated grad ──────────────────
198 let mut u_rot = vec![0.0f32; m * n];
199 for k in 0..m * n {
200 let g = g_rot[k] as f64;
201 let mi = b1 * st.m_rot[k] as f64 + (1.0 - b1) * g;
202 let vi = b2 * st.v_rot[k] as f64 + (1.0 - b2) * g * g;
203 st.m_rot[k] = mi as f32;
204 st.v_rot[k] = vi as f32;
205 let m_hat = mi / bc1;
206 let v_hat = vi / bc2;
207 u_rot[k] = (m_hat / (v_hat.sqrt() + eps as f64)) as f32;
208 }
209
210 // ── 5. Rotate back: U = Qₗ · U' · Q_rᵀ ───────────────────
211 // tmp = Qₗ · U'
212 matmul(&st.ql, &u_rot, m, m, n, &mut tmp);
213 // U = tmp · Q_rᵀ ⇒ U[i,j] = sum_p tmp[i,p] · Q_r[j,p]
214 let mut u = vec![0.0f32; m * n];
215 for i in 0..m {
216 for j in 0..n {
217 let mut s = 0.0f32;
218 for p in 0..n {
219 s += tmp[i * n + p] * st.qr[j * n + p];
220 }
221 u[i * n + j] = s;
222 }
223 }
224
225 // ── 6. Decoupled weight decay + parameter update ─────────
226 for i in 0..m * n {
227 param[i] -= lr * (u[i] + wd * param[i]);
228 }
229 }
230
231 fn end_iteration(&mut self) {
232 self.step += 1;
233 }
234}
235
236fn identity(n: usize) -> Vec<f32> {
237 let mut out = vec![0.0; n * n];
238 for i in 0..n {
239 out[i * n + i] = 1.0;
240 }
241 out
242}
243
244// Plain AdamW for the non-matrix fallback path.
245fn adamw_fallback(opt: &mut Soap, name: &str, param: &mut [f32], grad: &[f32]) {
246 let t = (opt.step + 1) as f64;
247 let b1 = opt.beta1 as f64;
248 let b2 = opt.beta2 as f64;
249 let bc1 = 1.0 - b1.powf(t);
250 let bc2 = 1.0 - b2.powf(t);
251 let m = zeros_entry(&mut opt.fb_m, name, param.len());
252 let v = zeros_entry(&mut opt.fb_v, name, param.len());
253 let eps = opt.eps as f64;
254 let lr = opt.lr as f64;
255 let wd = opt.weight_decay as f64;
256 for i in 0..param.len() {
257 let g = grad[i] as f64;
258 let mi = b1 * m[i] as f64 + (1.0 - b1) * g;
259 let vi = b2 * v[i] as f64 + (1.0 - b2) * g * g;
260 m[i] = mi as f32;
261 v[i] = vi as f32;
262 let p = param[i] as f64;
263 param[i] = (p - lr * (mi / bc1 / ((vi / bc2).sqrt() + eps) + wd * p)) as f32;
264 }
265}