Skip to main content

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}