Skip to main content

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}