rlx_optim/kron_psgd.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//! Kron-PSGD — Preconditioned SGD with a Kronecker-factored
6//! preconditioner (Li, 2018; "Preconditioned Stochastic Gradient
7//! Descent").
8//!
9//! # Idea
10//!
11//! Approximate the inverse Hessian as a Kronecker product
12//! `P ≈ P_L ⊗ P_R` where `P_L = Q_LᵀQ_L` and `P_R = Q_RᵀQ_R` for two
13//! upper-triangular factors. The factors are updated by a *Lie-group
14//! descent* on a whitening criterion — no eigendecomposition needed,
15//! and updates are stable by construction (the upper-triangular
16//! manifold).
17//!
18//! # Update rule
19//!
20//! For a 2-D parameter `W ∈ ℝ^{m×n}`:
21//!
22//! ```text
23//! A = Q_L · G · Q_Rᵀ // m×n
24//! B = Q_L⁻ᵀ · G · Q_R⁻¹ // m×n (triangular solves)
25//! dQ_L ∝ tril(A·Aᵀ − B·Bᵀ); Q_L ← Q_L − η_p · Q_L · dQ_L
26//! dQ_R ∝ tril(Aᵀ·A − Bᵀ·B); Q_R ← Q_R − η_p · Q_R · dQ_R
27//! P_L = Q_LᵀQ_L; P_R = Q_RᵀQ_R
28//! p_g = P_L · G · P_R // preconditioned grad
29//! [spectral-clip to ‖·‖_∞ ≤ clip, then SGD+momentum on p_g]
30//! ```
31//!
32//! Li (2018) Algorithm 1 uses an HVP probe `v` and its perturbed
33//! gradient to update `Q_L, Q_R`. This crate has no HVP oracle, so we
34//! use the gradient itself as the probe — the "PSGD-Affine"
35//! approximation — which is cheap and still gives strong empirical
36//! preconditioning on convex and mildly non-convex problems.
37//! Non-2-D parameters fall back to plain SGD-with-momentum.
38//!
39//! # When to use
40//!
41//! Ill-conditioned problems where Adam's coordinate-wise
42//! preconditioner is too weak (RNNs, deep MLPs, certain inverse
43//! problems). State cost per matrix: `m² + n²` plus a velocity buffer.
44
45use std::collections::HashMap;
46
47use crate::Optimizer;
48use crate::common::{matmul, zeros_entry};
49
50#[derive(Debug, Clone)]
51struct KronState {
52 ql: Vec<f32>, // m × m upper-triangular
53 qr: Vec<f32>, // n × n upper-triangular
54}
55
56/// Kron-PSGD — Kronecker-factored preconditioned SGD.
57#[derive(Debug, Clone)]
58pub struct KronPsgd {
59 /// Learning rate.
60 pub lr: f32,
61 /// Learning rate for the **preconditioner** update (Lie-group
62 /// descent on Q_L / Q_R). Default `0.1`. Too high ⇒ Q drifts;
63 /// too low ⇒ preconditioner lags.
64 pub precond_lr: f32,
65 /// Polyak momentum for the preconditioned-gradient SGD step.
66 /// Default `0.9`.
67 pub momentum: f32,
68 /// L2 weight-decay coefficient (folded into the gradient).
69 /// Default `0.0`.
70 pub weight_decay: f32,
71 /// Numerical floor on the preconditioner-update normalizer.
72 /// Default `1e-8`.
73 pub eps: f32,
74 /// Cap the per-coordinate magnitude of the preconditioned update
75 /// (defensive — early Q estimates can be ill-conditioned). Default `1.0`.
76 pub clip: f32,
77 state: HashMap<String, KronState>,
78 mom: HashMap<String, Vec<f32>>,
79}
80
81impl KronPsgd {
82 /// Construct with `(precond_lr, μ, λ, ε, clip) = (0.1, 0.9, 0.0, 1e-8, 1.0)`.
83 pub fn new(lr: f32) -> Self {
84 Self {
85 lr,
86 precond_lr: 0.1,
87 momentum: 0.9,
88 weight_decay: 0.0,
89 eps: 1e-8,
90 clip: 1.0,
91 state: HashMap::new(),
92 mom: HashMap::new(),
93 }
94 }
95
96 /// Override the Polyak momentum.
97 pub fn with_momentum(mut self, mu: f32) -> Self {
98 self.momentum = mu;
99 self
100 }
101
102 /// Override the weight-decay coefficient.
103 pub fn with_weight_decay(mut self, wd: f32) -> Self {
104 self.weight_decay = wd;
105 self
106 }
107}
108
109impl Optimizer for KronPsgd {
110 fn set_lr(&mut self, lr: f32) {
111 self.lr = lr;
112 }
113
114 fn step(&mut self, name: &str, shape: &[usize], param: &mut [f32], grad: &[f32]) {
115 debug_assert_eq!(param.len(), grad.len());
116 let lr = self.lr;
117 let wd = self.weight_decay;
118
119 if shape.len() != 2 {
120 // Non-matrix: SGD + momentum fallback.
121 let v = zeros_entry(&mut self.mom, name, param.len());
122 let mu = self.momentum;
123 for i in 0..param.len() {
124 v[i] = mu * v[i] + grad[i] + wd * param[i];
125 param[i] -= lr * v[i];
126 }
127 return;
128 }
129 let (m, n) = (shape[0], shape[1]);
130 debug_assert_eq!(m * n, param.len());
131 let st = self
132 .state
133 .entry(name.to_owned())
134 .or_insert_with(|| KronState {
135 ql: identity_triangular(m),
136 qr: identity_triangular(n),
137 });
138
139 // ── 1. Update Q_L, Q_R via Li (2018) Lie-group rule. ──────
140 // Use g itself as the probe; the affine variant requires:
141 // A = Q_L · g · Q_Rᵀ (m × n)
142 // B = Q_L⁻ᵀ · g · Q_R⁻¹ (m × n; cheap because Q is triangular)
143 // dQ_L ∝ tril(A·Aᵀ − B·Bᵀ); dQ_R ∝ tril(Aᵀ·A − Bᵀ·B).
144 let a = matmul_3(&st.ql, grad, &st.qr, m, n, /*trans_q_r=*/ true);
145 let b = matmul_3_inv(&st.ql, grad, &st.qr, m, n);
146 update_factor(&mut st.ql, &a, &b, m, n, true, self.precond_lr, self.eps);
147 update_factor(&mut st.qr, &a, &b, m, n, false, self.precond_lr, self.eps);
148
149 // ── 2. Preconditioned gradient: p_g = Q_Lᵀ · Q_L · g · Q_R · Q_Rᵀ ──
150 // Build Q_Lᵀ Q_L (m×m, symmetric)
151 let mut ql_t_ql = vec![0.0f32; m * m];
152 for i in 0..m {
153 for j in 0..m {
154 let mut s = 0.0f32;
155 for p in 0..m {
156 s += st.ql[p * m + i] * st.ql[p * m + j];
157 }
158 ql_t_ql[i * m + j] = s;
159 }
160 }
161 let mut qr_qr_t = vec![0.0f32; n * n];
162 for i in 0..n {
163 for j in 0..n {
164 let mut s = 0.0f32;
165 for p in 0..n {
166 s += st.qr[i * n + p] * st.qr[j * n + p];
167 }
168 qr_qr_t[i * n + j] = s;
169 }
170 }
171 // p_g = (Q_Lᵀ Q_L) · g · (Q_R Q_Rᵀ)
172 let mut tmp = vec![0.0f32; m * n];
173 matmul(&ql_t_ql, grad, m, m, n, &mut tmp);
174 let mut p_g = vec![0.0f32; m * n];
175 matmul(&tmp, &qr_qr_t, m, n, n, &mut p_g);
176
177 // ── 3. Spectral clip + momentum + apply. ─────────────────
178 let mut max_abs = 0.0f32;
179 for &x in &p_g {
180 if x.abs() > max_abs {
181 max_abs = x.abs();
182 }
183 }
184 let scale = if max_abs > self.clip {
185 self.clip / max_abs
186 } else {
187 1.0
188 };
189 let v = zeros_entry(&mut self.mom, name, param.len());
190 let mu = self.momentum;
191 for i in 0..param.len() {
192 let g = scale * p_g[i] + wd * param[i];
193 v[i] = mu * v[i] + g;
194 param[i] -= lr * v[i];
195 }
196 }
197}
198
199fn identity_triangular(n: usize) -> Vec<f32> {
200 let mut out = vec![0.0; n * n];
201 for i in 0..n {
202 out[i * n + i] = 1.0;
203 }
204 out
205}
206
207/// Compute `Q_L · G · Q_Rᵀ` (or `Q_L · G · Q_R` if `trans_q_r=false`).
208fn matmul_3(ql: &[f32], g: &[f32], qr: &[f32], m: usize, n: usize, trans_q_r: bool) -> Vec<f32> {
209 let mut t1 = vec![0.0f32; m * n];
210 matmul(ql, g, m, m, n, &mut t1);
211 let mut out = vec![0.0f32; m * n];
212 if trans_q_r {
213 // out = t1 · Q_Rᵀ ⇒ out[i,j] = sum_p t1[i,p] · Q_R[j,p]
214 for i in 0..m {
215 for j in 0..n {
216 let mut s = 0.0f32;
217 for p in 0..n {
218 s += t1[i * n + p] * qr[j * n + p];
219 }
220 out[i * n + j] = s;
221 }
222 }
223 } else {
224 matmul(&t1, qr, m, n, n, &mut out);
225 }
226 out
227}
228
229/// Compute `Q_L⁻ᵀ · G · Q_R⁻¹` for upper-triangular Q's via two
230/// triangular solves on `G`.
231fn matmul_3_inv(ql: &[f32], g: &[f32], qr: &[f32], m: usize, n: usize) -> Vec<f32> {
232 // First solve Q_Lᵀ · X = G column-by-column. Q_Lᵀ is lower-triangular.
233 let mut x = g.to_vec();
234 for j in 0..n {
235 // Forward-substitute one column.
236 for i in 0..m {
237 let mut s = x[i * n + j];
238 for p in 0..i {
239 s -= ql[p * m + i] * x[p * n + j];
240 }
241 let d = ql[i * m + i];
242 x[i * n + j] = if d.abs() > 1e-12 { s / d } else { 0.0 };
243 }
244 }
245 // Then solve Y · Q_R = X for Y row-by-row (Q_R upper-triangular).
246 // Equivalently: for each row i, back-substitute Y[i,:] · Q_R = X[i,:].
247 let mut y = x;
248 for i in 0..m {
249 for j in 0..n {
250 let mut s = y[i * n + j];
251 for p in 0..j {
252 s -= y[i * n + p] * qr[p * n + j];
253 }
254 let d = qr[j * n + j];
255 y[i * n + j] = if d.abs() > 1e-12 { s / d } else { 0.0 };
256 }
257 }
258 y
259}
260
261/// Lie-group update of a triangular factor. `which = true` updates Q_L
262/// using `A·Aᵀ − B·Bᵀ` (m×m), `which = false` updates Q_R using
263/// `Aᵀ·A − Bᵀ·B` (n×n). The descent direction is then projected onto
264/// the upper-triangular tangent space.
265fn update_factor(
266 q: &mut [f32],
267 a: &[f32],
268 b: &[f32],
269 m: usize,
270 n: usize,
271 which: bool,
272 plr: f32,
273 eps: f32,
274) {
275 let dim = if which { m } else { n };
276 let mut grad_q = vec![0.0f32; dim * dim];
277 // Build A·Aᵀ − B·Bᵀ or Aᵀ·A − Bᵀ·B.
278 let mut norm = 0.0f64;
279 for i in 0..dim {
280 for j in 0..dim {
281 let mut a_term = 0.0f32;
282 let mut b_term = 0.0f32;
283 if which {
284 for p in 0..n {
285 a_term += a[i * n + p] * a[j * n + p];
286 b_term += b[i * n + p] * b[j * n + p];
287 }
288 } else {
289 for p in 0..m {
290 a_term += a[p * n + i] * a[p * n + j];
291 b_term += b[p * n + i] * b[p * n + j];
292 }
293 }
294 let d = a_term - b_term;
295 grad_q[i * dim + j] = d;
296 norm += d as f64 * d as f64;
297 }
298 }
299 let scale = plr / ((norm.sqrt() as f32) + eps);
300 // Project onto upper-triangular: Q ← Q · (I − 0.5·scale·tril(grad_q + grad_qᵀ))
301 // (Simplified Lie-group projection; full version solves a tiny matrix
302 // exponential, but a single linearized step is the standard choice.)
303 for i in 0..dim {
304 for j in 0..dim {
305 if j < i {
306 grad_q[i * dim + j] = 0.0; // upper-triangular projection
307 }
308 }
309 }
310 let mut q_new = vec![0.0f32; dim * dim];
311 matmul(q, &grad_q, dim, dim, dim, &mut q_new);
312 for k in 0..dim * dim {
313 q[k] -= scale * q_new[k];
314 }
315}