Skip to main content

rlx_optim/
stiefel.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//! Stiefel — Riemannian SGD-with-momentum on the (compact) Stiefel
6//! manifold `St(m, n) = { W ∈ ℝ^{m×n} : W·Wᵀ = I_m }` (row-orthonormal,
7//! `m ≤ n`).
8//!
9//! # Why
10//!
11//! SPDNet's BiMap layer weights are constrained to be semi-orthogonal
12//! (`W·Wᵀ = I_m`). A plain Euclidean optimizer ignores that constraint
13//! and lets `W` drift off the manifold within a few steps, so the layer
14//! stops being a valid bilinear map on SPD matrices. This optimizer
15//! keeps every step *on* the manifold: it projects the Euclidean
16//! gradient to the tangent space, moves along the tangent, then
17//! **retracts** back to `St(m, n)`.
18//!
19//! # Update rule (2-D parameter `W ∈ ℝ^{m×n}`, `m ≤ n`)
20//!
21//! ```text
22//! sym(A) = ½·(A + Aᵀ)
23//! G_riem = G − W · sym(Wᵀ·G)          // project Euclidean grad to T_W St
24//! v_t    = μ·v_{t-1} + G_riem          // (optional) tangent momentum
25//! Y      = W − lr · v_t                // Euclidean step off the manifold
26//! W_new  = qf(Y)                        // QR retraction back onto St(m,n)
27//! ```
28//!
29//! where `qf(Y)` is the row-orthonormal factor of `Y`: we take the thin
30//! QR of `Yᵀ` (shape `n×m`, tall) and return `Qᵀ` (shape `m×n`) so that
31//! `W_new·W_newᵀ = I_m`. The QR is a dependency-free modified
32//! Gram–Schmidt on the rows of `Y`. The sign convention of `qf`
33//! (diagonal of `R` forced positive) keeps the retraction a smooth map
34//! agreeing with the exponential to first order — the standard QR
35//! retraction of Absil, Mahony & Sepulchre, *Optimization Algorithms on
36//! Matrix Manifolds* (Princeton, 2008), §4.1.
37//!
38//! # Non-Stiefel parameters
39//!
40//! Only 2-D parameters with `rows ≤ cols` are treated as Stiefel
41//! points. Any other shape (bias vectors, `rows > cols`, higher-rank
42//! tensors) falls back to plain SGD-with-momentum so a single `Stiefel`
43//! instance can still drive an entire SPDNet (BiMap weights on the
44//! manifold, everything else Euclidean).
45//!
46//! # State
47//!
48//! One tangent-momentum buffer per parameter (only allocated / used when
49//! `momentum > 0`), same footprint as [`crate::Sgd`].
50
51use std::collections::HashMap;
52
53use crate::Optimizer;
54use crate::common::zeros_entry;
55
56/// Riemannian SGD-with-momentum on the Stiefel manifold.
57///
58/// All hyperparameters are public so callers can hot-swap them between
59/// iterations. State is keyed by parameter name.
60#[derive(Debug, Clone)]
61pub struct Stiefel {
62    /// Learning rate (tangent step size). No default — pass to
63    /// [`Stiefel::new`].
64    pub lr: f32,
65    /// Polyak momentum coefficient μ ∈ \[0, 1). `0.0` disables the
66    /// tangent-momentum buffer entirely. Default `0.0`.
67    pub momentum: f32,
68    m: HashMap<String, Vec<f32>>,
69}
70
71impl Stiefel {
72    /// Construct with the given learning rate and momentum disabled.
73    pub fn new(lr: f32) -> Self {
74        Self {
75            lr,
76            momentum: 0.0,
77            m: HashMap::new(),
78        }
79    }
80
81    /// Enable Polyak momentum on the tangent update.
82    pub fn with_momentum(mut self, momentum: f32) -> Self {
83        self.momentum = momentum;
84        self
85    }
86}
87
88impl Optimizer for Stiefel {
89    fn set_lr(&mut self, lr: f32) {
90        self.lr = lr;
91    }
92
93    fn step(&mut self, name: &str, shape: &[usize], param: &mut [f32], grad: &[f32]) {
94        debug_assert_eq!(param.len(), grad.len());
95        let lr = self.lr;
96        let mu = self.momentum;
97
98        // Non-Stiefel shapes: plain SGD-with-momentum fallback.
99        let is_stiefel = shape.len() == 2 && shape[0] <= shape[1];
100        if !is_stiefel {
101            if mu == 0.0 {
102                for i in 0..param.len() {
103                    param[i] -= lr * grad[i];
104                }
105            } else {
106                let m = zeros_entry(&mut self.m, name, param.len());
107                for i in 0..param.len() {
108                    m[i] = mu * m[i] + grad[i];
109                    param[i] -= lr * m[i];
110                }
111            }
112            return;
113        }
114
115        let (rows, cols) = (shape[0], shape[1]);
116        debug_assert_eq!(rows * cols, param.len());
117
118        // ── Riemannian gradient ─────────────────────────────────────
119        // sym = ½(WᵀG + GᵀW) is m×m via Wᵀ·G symmetrized; but we need
120        // W·sym(Wᵀ·G) which is m×n. Compute S = Wᵀ·G? No: for the
121        // row-orthonormal convention the projection is
122        //   G_riem = G − W · sym(Wᵀ·G_no)   where sym acts on the m×m
123        // matrix  M = W·Gᵀ  (since Wᵀ·G would be n×n).
124        //
125        // Tangent space of St(m,n) at W (rows orthonormal):
126        //   T_W = { Z : W·Zᵀ + Z·Wᵀ = 0 }.
127        // The orthogonal projection of an ambient G onto T_W is
128        //   P_W(G) = G − W · sym(W·Gᵀ),   sym(A)=½(A+Aᵀ),  A = W·Gᵀ (m×m).
129        // (Edelman, Arias & Smith 1998; ASM 2008 §3.6.1, transposed to
130        // the wide/row-orthonormal layout.)
131        let a = mat_mul_bt(param, grad, rows, cols); // A = W·Gᵀ  (m×m)
132        let mut symm = vec![0.0f32; rows * rows];
133        for i in 0..rows {
134            for j in 0..rows {
135                symm[i * rows + j] = 0.5 * (a[i * rows + j] + a[j * rows + i]);
136            }
137        }
138        // W·sym  (m×n)
139        let wsym = mat_mul(&symm, param, rows, rows, cols);
140        let mut g_riem = vec![0.0f32; rows * cols];
141        for i in 0..rows * cols {
142            g_riem[i] = grad[i] - wsym[i];
143        }
144
145        // ── Tangent step (+ optional momentum) ──────────────────────
146        let mut y = vec![0.0f32; rows * cols];
147        if mu == 0.0 {
148            for i in 0..rows * cols {
149                y[i] = param[i] - lr * g_riem[i];
150            }
151        } else {
152            let m = zeros_entry(&mut self.m, name, rows * cols);
153            for i in 0..rows * cols {
154                m[i] = mu * m[i] + g_riem[i];
155                y[i] = param[i] - lr * m[i];
156            }
157        }
158
159        // ── QR retraction back onto St(m,n) ─────────────────────────
160        // qf(Y) = row-orthonormal factor. Orthonormalize the rows of Y
161        // (modified Gram–Schmidt); result satisfies W_new·W_newᵀ = I_m.
162        qf_rows_in_place(&mut y, rows, cols);
163        param.copy_from_slice(&y);
164    }
165}
166
167/// Row-major `A · B` with `A: m×k`, `B: k×n` → `m×n`. Local copy so
168/// this module stays self-contained (the crate's `common::matmul`
169/// writes into a caller buffer; here an owned return reads cleaner).
170fn mat_mul(a: &[f32], b: &[f32], m: usize, k: usize, n: usize) -> Vec<f32> {
171    debug_assert_eq!(a.len(), m * k);
172    debug_assert_eq!(b.len(), k * n);
173    let mut c = vec![0.0f32; m * n];
174    for i in 0..m {
175        for j in 0..n {
176            let mut s = 0.0f32;
177            for p in 0..k {
178                s += a[i * k + p] * b[p * n + j];
179            }
180            c[i * n + j] = s;
181        }
182    }
183    c
184}
185
186/// Row-major `A · Bᵀ` with `A: m×n`, `B: r×n` → `m×r`.
187fn mat_mul_bt(a: &[f32], b: &[f32], m: usize, n: usize) -> Vec<f32> {
188    // Here both A and B are m×n (W and G); result is m×m.
189    debug_assert_eq!(a.len(), m * n);
190    debug_assert_eq!(b.len(), m * n);
191    let mut c = vec![0.0f32; m * m];
192    for i in 0..m {
193        for j in 0..m {
194            let mut s = 0.0f32;
195            for p in 0..n {
196                s += a[i * n + p] * b[j * n + p];
197            }
198            c[i * m + j] = s;
199        }
200    }
201    c
202}
203
204/// In-place `qf` retraction: orthonormalize the `rows` rows of the
205/// row-major `rows × cols` matrix `y` (each row a length-`cols`
206/// vector) via **modified Gram–Schmidt**, so afterwards
207/// `Y·Yᵀ = I_rows`. This is exactly the thin-QR `Q` factor of `Yᵀ`
208/// transposed back, with the diagonal of `R` implicitly forced
209/// non-negative (each pivot row is normalized by its own positive
210/// norm), matching the standard QR-retraction sign convention. Requires
211/// `rows ≤ cols` and the rows of `Y` to be (numerically) linearly
212/// independent — true for any `Y` close to a Stiefel point.
213fn qf_rows_in_place(y: &mut [f32], rows: usize, cols: usize) {
214    debug_assert_eq!(y.len(), rows * cols);
215    for i in 0..rows {
216        // Subtract projections onto the already-orthonormal rows 0..i.
217        for j in 0..i {
218            let mut dot = 0.0f32;
219            for c in 0..cols {
220                dot += y[i * cols + c] * y[j * cols + c];
221            }
222            for c in 0..cols {
223                y[i * cols + c] -= dot * y[j * cols + c];
224            }
225        }
226        // Normalize row i.
227        let mut nrm = 0.0f64;
228        for c in 0..cols {
229            let v = y[i * cols + c] as f64;
230            nrm += v * v;
231        }
232        let nrm = (nrm.sqrt() as f32).max(1e-12);
233        let inv = 1.0 / nrm;
234        for c in 0..cols {
235            y[i * cols + c] *= inv;
236        }
237    }
238}