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}