rlx_optim/muon.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//! Muon — MomentUm Orthogonalized by Newton–Schulz (Jordan, Bernstein,
6//! Vyas, Hubara, et al., 2024).
7//!
8//! # Idea
9//!
10//! For a 2-D parameter, replace the momentum buffer with its **closest
11//! semi-orthogonal matrix** before applying it as an update. The SVD
12//! `M = U·Σ·Vᵀ` has closest semi-orthogonal matrix `U·Vᵀ` — but the
13//! SVD is expensive. A *Newton–Schulz cubic iteration* approximates
14//! `U·Vᵀ` in only 5 small matrix products per step. Empirically this
15//! gives a step-size-invariant update that punches above its weight on
16//! transformer training.
17//!
18//! # Update rule (2-D parameter `W ∈ ℝ^{m×n}`)
19//!
20//! ```text
21//! m_t = μ·m_{t-1} + g_t // Polyak momentum
22//! M = m_t if !nesterov
23//! = g_t + μ·m_t if nesterov
24//! M̂ = M / ‖M‖_F // normalize for NS
25//! repeat ns_steps times: // ns_steps = 5
26//! A = M̂ · M̂ᵀ
27//! M̂ ← a·M̂ + b·A·M̂ + c·A²·M̂ // cubic NS iter
28//! U = √max(m, n) · M̂ // RMS-of-cols scaling
29//! θ_t = θ_{t-1} − lr · ( U + λ·θ_{t-1} )
30//! ```
31//!
32//! The (a, b, c) coefficients are chosen so the cubic polynomial maps
33//! singular values in (0, √3] toward 1; defaults
34//! `(3.4445, −4.7750, 2.0315)` are from the original release.
35//!
36//! Non-2-D parameters fall back to SGD-with-momentum (the original
37//! recipe routes them to AdamW; this crate stays dependency-free).
38//!
39//! # When to use
40//!
41//! Pre-training transformer matrix-shaped weights (Q/K/V/FFN
42//! projections). Often paired with AdamW for embeddings and biases.
43//! State cost: one momentum buffer per matrix.
44
45use std::collections::HashMap;
46
47use crate::Optimizer;
48use crate::common::zeros_entry;
49
50/// Muon — Momentum-Orthogonalized-by-Newton-Schulz.
51///
52/// Per-tensor state: **one** momentum buffer per matrix (half of
53/// Adam's footprint, like Lion).
54#[derive(Debug, Clone)]
55pub struct Muon {
56 /// Learning rate. The Newton–Schulz update has roughly unit
57 /// Frobenius norm per column, so this is on the same scale as
58 /// SGD's lr — typically `2e-2` to `5e-2`.
59 pub lr: f32,
60 /// Polyak momentum coefficient. Default `0.95`.
61 pub momentum: f32,
62 /// Use Nesterov lookahead inside the matrix being orthogonalized.
63 /// Default `true`.
64 pub nesterov: bool,
65 /// Decoupled weight-decay coefficient λ. Default `0.0`.
66 pub weight_decay: f32,
67 /// Newton–Schulz iteration count. `5` is the published default;
68 /// `3` is enough for most well-conditioned matrices.
69 pub ns_steps: u32,
70 /// `(a, b, c)` coefficients of the cubic Newton–Schulz iteration
71 /// `X ← a·X + b·(XXᵀ)X + c·(XXᵀ)²X`. Defaults match Jordan et al.
72 pub ns_coeffs: (f32, f32, f32),
73 m: HashMap<String, Vec<f32>>,
74}
75
76impl Muon {
77 /// Construct with `(μ, nesterov, λ, ns_steps) = (0.95, true, 0.0, 5)`
78 /// and the published NS coefficients.
79 pub fn new(lr: f32) -> Self {
80 Self {
81 lr,
82 momentum: 0.95,
83 nesterov: true,
84 weight_decay: 0.0,
85 ns_steps: 5,
86 ns_coeffs: (3.4445, -4.7750, 2.0315),
87 m: HashMap::new(),
88 }
89 }
90
91 /// Override the Polyak momentum coefficient.
92 pub fn with_momentum(mut self, mu: f32) -> Self {
93 self.momentum = mu;
94 self
95 }
96
97 /// Override the decoupled-decay coefficient.
98 pub fn with_weight_decay(mut self, wd: f32) -> Self {
99 self.weight_decay = wd;
100 self
101 }
102
103 /// Override the Newton–Schulz iteration count.
104 pub fn with_ns_steps(mut self, n: u32) -> Self {
105 self.ns_steps = n;
106 self
107 }
108}
109
110impl Optimizer for Muon {
111 fn set_lr(&mut self, lr: f32) {
112 self.lr = lr;
113 }
114
115 fn step(&mut self, name: &str, shape: &[usize], param: &mut [f32], grad: &[f32]) {
116 debug_assert_eq!(param.len(), grad.len());
117 let mu = self.momentum;
118 let wd = self.weight_decay;
119 let lr = self.lr;
120 let m = zeros_entry(&mut self.m, name, param.len());
121 // EMA buffer (classical Polyak momentum: `m ← μ·m + g`).
122 for i in 0..param.len() {
123 m[i] = mu * m[i] + grad[i];
124 }
125 if shape.len() != 2 {
126 // Non-matrix: SGD-with-momentum update.
127 for i in 0..param.len() {
128 let g = if self.nesterov {
129 grad[i] + mu * m[i]
130 } else {
131 m[i]
132 };
133 param[i] -= lr * (g + wd * param[i]);
134 }
135 return;
136 }
137 let (rows, cols) = (shape[0], shape[1]);
138 debug_assert_eq!(rows * cols, param.len());
139 // Build the matrix to orthogonalize. With Nesterov:
140 // G = grad + μ·m (m has already been updated above)
141 let mut g_mat = vec![0.0f32; rows * cols];
142 if self.nesterov {
143 for i in 0..rows * cols {
144 g_mat[i] = grad[i] + mu * m[i];
145 }
146 } else {
147 g_mat.copy_from_slice(m);
148 }
149 let ortho = newton_schulz_orth(&g_mat, rows, cols, self.ns_steps, self.ns_coeffs);
150 // The Muon paper scales the update by sqrt(max(rows, cols)) so
151 // its effective magnitude matches a unit-norm column.
152 let s = (rows.max(cols) as f32).sqrt();
153 for i in 0..param.len() {
154 param[i] -= lr * (s * ortho[i] + wd * param[i]);
155 }
156 }
157}
158
159/// Map `f(row_index, &mut row)` over the `n_rows × cols` row-major matrix
160/// `out`, parallelized across output rows. With the `parallel` feature this uses
161/// rayon's **persistent global pool** (no per-call thread spawn — the dominant
162/// overhead when the Newton–Schulz iteration fires this many small matmuls per
163/// step); without it, dependency-free scoped threads. Serial below 64 rows.
164/// `f` reads only shared (immutable) state and writes its own disjoint row.
165/// (Dead on macOS, where the Newton–Schulz path uses Accelerate BLAS instead.)
166#[allow(dead_code)]
167fn par_rows<F: Fn(usize, &mut [f32]) + Sync>(out: &mut [f32], cols: usize, f: F) {
168 let n = out.len() / cols;
169 if n < 64 {
170 for (i, row) in out.chunks_mut(cols).enumerate() {
171 f(i, row);
172 }
173 return;
174 }
175 #[cfg(feature = "parallel")]
176 {
177 use rayon::prelude::*;
178 out.par_chunks_mut(cols)
179 .enumerate()
180 .for_each(|(i, row)| f(i, row));
181 }
182 #[cfg(not(feature = "parallel"))]
183 {
184 let threads = std::thread::available_parallelism()
185 .map(|x| x.get())
186 .unwrap_or(1)
187 .min(n);
188 let per = n.div_ceil(threads);
189 std::thread::scope(|s| {
190 for (t, chunk) in out.chunks_mut(per * cols).enumerate() {
191 let f = &f;
192 s.spawn(move || {
193 let base = t * per;
194 for (ii, row) in chunk.chunks_mut(cols).enumerate() {
195 f(base + ii, row);
196 }
197 });
198 }
199 });
200 }
201}
202
203/// Accelerate (Apple AMX/SME) BLAS shims for the Newton–Schulz matmuls. Linked
204/// via `build.rs` (`framework=Accelerate`). Row-major `cblas_sgemm`; the weight
205/// matrices are never empty so no zero-dim guard is needed.
206#[cfg(target_os = "macos")]
207#[allow(unsafe_code)]
208mod accel {
209 unsafe extern "C" {
210 #[link_name = "cblas_sgemm"]
211 fn cblas_sgemm(
212 order: i32,
213 transa: i32,
214 transb: i32,
215 m: i32,
216 n: i32,
217 k: i32,
218 alpha: f32,
219 a: *const f32,
220 lda: i32,
221 b: *const f32,
222 ldb: i32,
223 beta: f32,
224 c: *mut f32,
225 ldc: i32,
226 );
227 }
228 const ROW_MAJOR: i32 = 101;
229 const NO_TRANS: i32 = 111;
230 const TRANS: i32 = 112;
231
232 /// `C[m×n] = A[m×k] · B[k×n]`.
233 #[inline]
234 pub fn gemm_nn(a: &[f32], b: &[f32], c: &mut [f32], m: usize, k: usize, n: usize) {
235 unsafe {
236 cblas_sgemm(
237 ROW_MAJOR,
238 NO_TRANS,
239 NO_TRANS,
240 m as i32,
241 n as i32,
242 k as i32,
243 1.0,
244 a.as_ptr(),
245 k as i32,
246 b.as_ptr(),
247 n as i32,
248 0.0,
249 c.as_mut_ptr(),
250 n as i32,
251 );
252 }
253 }
254
255 /// `C[m×n] = A[m×k] · B[n×k]ᵀ`.
256 #[inline]
257 pub fn gemm_nt(a: &[f32], b: &[f32], c: &mut [f32], m: usize, k: usize, n: usize) {
258 unsafe {
259 cblas_sgemm(
260 ROW_MAJOR,
261 NO_TRANS,
262 TRANS,
263 m as i32,
264 n as i32,
265 k as i32,
266 1.0,
267 a.as_ptr(),
268 k as i32,
269 b.as_ptr(),
270 k as i32,
271 0.0,
272 c.as_mut_ptr(),
273 n as i32,
274 );
275 }
276 }
277}
278
279/// Newton–Schulz semi-orthogonalization. Operates on a row-major
280/// `rows × cols` matrix and returns its closest semi-orthogonal matrix
281/// (up to the polynomial truncation). The input is first scaled by its
282/// Frobenius norm to stay inside the polynomial's region of convergence.
283/// The per-iteration `X·Xᵀ`, `A²` and `X`-update products are parallelized
284/// over rows, so it stays cheap even for the largest transformer matrices.
285pub fn newton_schulz_orth(
286 g: &[f32],
287 rows: usize,
288 cols: usize,
289 steps: u32,
290 c: (f32, f32, f32),
291) -> Vec<f32> {
292 let mut x = g.to_vec();
293 // Frobenius normalization.
294 let mut fro = 0.0f64;
295 for &xi in &x {
296 fro += xi as f64 * xi as f64;
297 }
298 let fro = (fro.sqrt() as f32).max(1e-12);
299 for xi in &mut x {
300 *xi /= fro;
301 }
302 // Operate with the SMALLER dimension as rows, so the `XXᵀ` Gram (`r×r`) is
303 // `min×min` — the cubic NS iteration's cost is dominated by that `r×r`
304 // matrix, and `U·Vᵀ` is identical in either orientation. (Transposing to
305 // `min×max`, as the canonical Muon reference does, is up to (max/min)²
306 // cheaper on the `A²` product than working on the `max×max` Gram.)
307 let (mut x_mat, r, k, transposed) = if rows > cols {
308 // transpose rows×cols → cols×rows so r = cols = min
309 let mut t = vec![0.0f32; rows * cols];
310 for i in 0..rows {
311 for j in 0..cols {
312 t[j * rows + i] = x[i * cols + j];
313 }
314 }
315 (t, cols, rows, true)
316 } else {
317 (x, rows, cols, false)
318 };
319 let (a, b, cc) = c;
320 let mut tmp = vec![0.0f32; r * k]; // XXᵀ X has shape r × k
321 let mut a_mat = vec![0.0f32; r * r];
322 let mut a2 = vec![0.0f32; r * r];
323 for _ in 0..steps {
324 // On macOS the three per-iteration products (`X·Xᵀ`, `A²`, and the
325 // `X`-update contraction) are the whole cost of Muon on large models —
326 // route them through Accelerate's `cblas_sgemm`, which dispatches to the
327 // AMX / SME matrix coprocessor and is ~10-50× the hand-rolled loop.
328 #[cfg(target_os = "macos")]
329 {
330 accel::gemm_nt(&x_mat, &x_mat, &mut a_mat, r, k, r); // A = X · Xᵀ
331 accel::gemm_nn(&a_mat, &a_mat, &mut a2, r, r, r); // A² = A · A
332 for i in 0..r * r {
333 a2[i] = b * a_mat[i] + cc * a2[i]; // reuse a2 as (b·A + cc·A²)
334 }
335 accel::gemm_nn(&a2, &x_mat, &mut tmp, r, r, k); // tmp = (b·A + cc·A²)·X
336 for i in 0..r * k {
337 tmp[i] += a * x_mat[i]; // tmp += a·X
338 }
339 }
340 // Portable parallel fallback (non-macOS): cache-friendly `ikj` matmuls.
341 #[cfg(not(target_os = "macos"))]
342 {
343 // A = X · Xᵀ — dot products of contiguous rows.
344 par_rows(&mut a_mat, r, |i, arow| {
345 let xi = &x_mat[i * k..i * k + k];
346 for (j, aij) in arow.iter_mut().enumerate() {
347 let xj = &x_mat[j * k..j * k + k];
348 let mut s = 0.0f32;
349 for p in 0..k {
350 s += xi[p] * xj[p];
351 }
352 *aij = s;
353 }
354 });
355 // A² = A · A.
356 par_rows(&mut a2, r, |i, a2row| {
357 a2row.fill(0.0);
358 let ai = &a_mat[i * r..i * r + r];
359 for p in 0..r {
360 let aip = ai[p];
361 let ap = &a_mat[p * r..p * r + r];
362 for j in 0..r {
363 a2row[j] += aip * ap[j];
364 }
365 }
366 });
367 // X ← a·X + b·A·X + cc·A²·X.
368 par_rows(&mut tmp, k, |i, trow| {
369 let xi = &x_mat[i * k..i * k + k];
370 for j in 0..k {
371 trow[j] = a * xi[j];
372 }
373 let ai = &a_mat[i * r..i * r + r];
374 let a2i = &a2[i * r..i * r + r];
375 for p in 0..r {
376 let coef = b * ai[p] + cc * a2i[p];
377 let xp = &x_mat[p * k..p * k + k];
378 for j in 0..k {
379 trow[j] += coef * xp[j];
380 }
381 }
382 });
383 }
384 std::mem::swap(&mut x_mat, &mut tmp);
385 }
386 if transposed {
387 // Transpose back to rows × cols.
388 let mut out = vec![0.0f32; rows * cols];
389 for i in 0..r {
390 for j in 0..k {
391 out[j * r + i] = x_mat[i * k + j];
392 }
393 }
394 out
395 } else {
396 x_mat
397 }
398}