Skip to main content

ruda_optim/optim/muon/
mod.rs

1
2use ruda_model::{module::AutodiffModule, record::Record};
3
4use ruda_model::config::Config;
5use ruda_model::tensor::{Tensor, backend::AutodiffBackend};
6use ruda_model::tensor::{backend::Backend, ops::Device};
7use serde::{Deserialize, Serialize};
8
9use super::{
10    SimpleOptimizer,
11    adaptor::OptimizerAdaptor,
12    decay::WeightDecayConfig,
13    momentum::{Momentum, MomentumConfig, MomentumState},
14};
15use crate::LearningRate;
16use ruda_model::tensor::DType;
17
18mod error;
19pub use error::MuonError;
20mod grouped;
21pub use grouped::{MuonAdamW, MuonAdamWConfig, MuonAdamWRecord};
22
23/// Momentum convention. Checkpoint buffers are NOT interchangeable between modes.
24#[derive(Clone, Default, Debug, Copy, PartialEq, Eq, Serialize, Deserialize)]
25pub enum MuonMomentumMode {
26    /// Preserve RUDA's existing SGD momentum (including first-step initialization).
27    #[default]
28    Sgd,
29    /// Exponential moving average: m = beta*m + (1-beta)*g, starting at zero.
30    /// Nesterov update is (1-beta)*g + beta*m. Dampening must be zero.
31    Ema,
32}
33
34#[cfg(not(feature = "std"))]
35#[allow(unused_imports)]
36use num_traits::Float as _;
37
38/// Logical orientation used ONLY for shape-based learning-rate scaling.
39/// The tensor itself is not reshaped and its update keeps the same geometry.
40#[derive(Clone, Default, Debug, Copy, PartialEq, Eq, Serialize, Deserialize)]
41pub enum MuonMatrixLayout {
42    /// Interpret [rows, columns] as [outputs, inputs]; legacy behavior.
43    #[default]
44    AsStored,
45    /// Tensor is [inputs, outputs], as in ruda-nn Linear. Use the reversed
46    /// aspect ratio for Original LR scaling (the RMS rule is symmetric).
47    InputOutput,
48}
49
50/// Learning rate adjustment method for Muon optimizer.
51///
52/// Muon adjusts the learning rate based on parameter shape to maintain consistent
53/// RMS across rectangular matrices.
54///
55/// # References
56///
57/// - Original: [Muon: An optimizer for hidden layers](https://kellerjordan.github.io/posts/muon/)
58/// - Moonshot: [Muon is Scalable for LLM Training](https://arxiv.org/pdf/2502.16982)
59#[derive(Clone, Default, Debug, Copy, PartialEq, Eq, Serialize, Deserialize)]
60pub enum AdjustLrFn {
61    /// Keller Jordan's original method: `lr * sqrt(max(1, A/B))`
62    ///
63    /// This scales the learning rate based on the aspect ratio of the weight matrix,
64    /// ensuring that tall matrices (more rows than columns) get proportionally larger
65    /// learning rates.
66    ///
67    /// # Example
68    ///
69    /// For a [1024, 512] matrix: `lr * sqrt(1024/512) = lr * 1.414`
70    #[default]
71    Original,
72
73    /// Moonshot's method: `lr * 0.2 * sqrt(max(A, B))`
74    ///
75    /// This method is designed to match AdamW's RMS, allowing Muon to directly reuse
76    /// learning rates and weight decay values tuned for AdamW without retuning.
77    ///
78    /// # Example
79    ///
80    /// For a [1024, 512] matrix: `lr * 0.2 * sqrt(1024) = lr * 6.4`
81    MatchRmsAdamW,
82}
83
84impl AdjustLrFn {
85    /// Calculate the learning rate adjustment ratio for a given parameter shape.
86    ///
87    /// # Arguments
88    ///
89    /// * `shape` - Parameter shape (uses first two dimensions)
90    ///
91    /// # Returns
92    ///
93    /// Adjustment ratio to multiply with the base learning rate
94    fn adjustment_ratio(&self, shape: &[usize]) -> f64 {
95        if shape.len() < 2 {
96            return 1.0;
97        }
98
99        let a = shape[0] as f64;
100        let b = shape[1] as f64;
101
102        match self {
103            Self::Original => {
104                // sqrt(max(1, A/B))
105                let ratio = a / b;
106                ratio.max(1.0).sqrt()
107            }
108            Self::MatchRmsAdamW => {
109                // 0.2 * sqrt(max(A, B))
110                0.2 * a.max(b).sqrt()
111            }
112        }
113    }
114}
115
116/// Muon configuration.
117///
118/// Muon is an optimizer specifically designed for 2D parameters of neural network
119/// hidden layers (weight matrices). Other parameters such as biases and embeddings
120/// should be optimized using a standard method such as AdamW.
121///
122/// # Learning Rate Adjustment
123///
124/// Muon adjusts the learning rate based on parameter shape to maintain consistent
125/// RMS across rectangular matrices. Two methods are available:
126///
127/// - **Original**: Uses `sqrt(max(1, A/B))` where A and B are the first two dimensions.
128///   This is Keller Jordan's method and is the default.
129///
130/// - **MatchRmsAdamW**: Uses `0.2 * sqrt(max(A, B))`. This is Moonshot's method
131///   designed to match AdamW's RMS, allowing direct reuse of AdamW hyperparameters.
132///
133/// # Example
134///
135/// ```ignore
136/// use ruda_optim::{MuonConfig, AdjustLrFn};
137///
138/// // Using default (Original) method
139/// let optimizer = MuonConfig::new().init();
140///
141/// // Using MatchRmsAdamW for AdamW-compatible hyperparameters
142/// let optimizer = MuonConfig::new()
143///     .with_adjust_lr_fn(AdjustLrFn::MatchRmsAdamW)
144///     .init();
145/// ```
146///
147/// # References
148///
149/// - [Muon: An optimizer for hidden layers in neural networks](https://kellerjordan.github.io/posts/muon/)
150/// - [Muon is Scalable for LLM Training](https://arxiv.org/pdf/2502.16982)
151/// - [PyTorch Implementation](https://github.com/pytorch/pytorch/blob/main/torch/optim/muon.py)
152/// - [Original Implementation](https://github.com/KellerJordan/Muon)
153#[derive(Config, Debug)]
154pub struct MuonConfig {
155    /// [Weight decay](WeightDecayConfig) config.
156    weight_decay: Option<WeightDecayConfig>,
157
158    /// [Momentum](MomentumConfig) config.
159    ///
160    /// Muon always uses momentum. Default configuration:
161    /// - momentum: 0.95
162    /// - dampening: 0.0
163    /// - nesterov: true
164    #[config(default = "MomentumConfig { momentum: 0.95, dampening: 0.0, nesterov: true }")]
165    momentum: MomentumConfig,
166
167    /// Newton-Schulz iteration coefficients (a, b, c).
168    ///
169    /// These coefficients are selected to maximize the slope at zero for the
170    /// quintic iteration. Default values are from Keller Jordan's implementation.
171    #[config(default = "(3.4445, -4.775, 2.0315)")]
172    ns_coefficients: (f32, f32, f32),
173
174    /// Epsilon for numerical stability.
175    #[config(default = 1e-7)]
176    epsilon: f32,
177
178    /// Number of Newton-Schulz iteration steps.
179    #[config(default = 5)]
180    ns_steps: usize,
181
182    /// Learning rate adjustment method.
183    ///
184    /// Controls how the learning rate is adjusted based on parameter shape.
185    /// See [`AdjustLrFn`] for available methods.
186    #[config(default = "AdjustLrFn::Original")]
187    adjust_lr_fn: AdjustLrFn,
188
189    /// Legacy SGD momentum by default; choose Ema explicitly for the current
190    /// reference implementation's buffer convention. Save this config with records.
191    #[config(default = "MuonMomentumMode::Sgd")]
192    momentum_mode: MuonMomentumMode,
193
194    /// Scale before squaring to avoid FP32 Frobenius-norm overflow/underflow.
195    /// Opt-in to preserve legacy rounding. Requires FP32 parameters and gradients.
196    #[config(default = false)]
197    stable_normalization: bool,
198
199    /// Explicit matrix orientation; embeddings are not identified by this flag.
200    #[config(default = "MuonMatrixLayout::AsStored")]
201    matrix_layout: MuonMatrixLayout,
202}
203
204impl MuonConfig {
205    /// Build a [`Muon`] from the config.
206    pub fn build<B: Backend>(&self) -> Muon<B> {
207        self.try_build().unwrap_or_else(|error| panic!("{error}"))
208    }
209
210    /// Validate host configuration without launching a device operation.
211    pub fn validate(&self) -> Result<(), MuonError> {
212        let beta = self.momentum.momentum;
213        let dampening = self.momentum.dampening;
214        if !beta.is_finite() || !(0.0..1.0).contains(&beta) {
215            return Err(MuonError::InvalidConfig("momentum must be finite in [0, 1)"));
216        }
217        if !dampening.is_finite() || !(0.0..=1.0).contains(&dampening) {
218            return Err(MuonError::InvalidConfig("dampening must be finite in [0, 1]"));
219        }
220        if self.momentum.nesterov && (beta == 0.0 || dampening != 0.0) {
221            return Err(MuonError::InvalidConfig("Nesterov requires positive momentum and zero dampening"));
222        }
223        if self.momentum_mode == MuonMomentumMode::Ema && dampening != 0.0 {
224            return Err(MuonError::InvalidConfig("EMA momentum requires zero dampening"));
225        }
226        if !self.epsilon.is_finite() || self.epsilon <= 0.0 {
227            return Err(MuonError::InvalidConfig("epsilon must be finite and positive"));
228        }
229        if !(1..100).contains(&self.ns_steps) {
230            return Err(MuonError::InvalidConfig("Newton-Schulz steps must be in 1..100"));
231        }
232        let (a, b, c) = self.ns_coefficients;
233        if !a.is_finite() || !b.is_finite() || !c.is_finite() {
234            return Err(MuonError::InvalidConfig("Newton-Schulz coefficients must be finite"));
235        }
236        if let Some(decay) = &self.weight_decay {
237            if !decay.penalty.is_finite() || decay.penalty < 0.0 {
238                return Err(MuonError::InvalidConfig("weight decay must be finite and nonnegative"));
239            }
240        }
241        Ok(())
242    }
243
244    /// Build with explicit configuration errors instead of a panic.
245    pub fn try_build<B: Backend>(&self) -> Result<Muon<B>, MuonError> {
246        self.validate()?;
247        let momentum = Momentum::new(&self.momentum);
248        let weight_decay_penalty = self.weight_decay.as_ref().map(|wd| wd.penalty);
249
250        Ok(Muon {
251            momentum,
252            ns_params: NewtonSchulzParams::new(self.ns_coefficients, self.ns_steps),
253            weight_decay_penalty,
254            epsilon: self.epsilon,
255            adjust_lr_fn: self.adjust_lr_fn,
256            momentum_mode: self.momentum_mode,
257            momentum_beta: self.momentum.momentum,
258            nesterov: self.momentum.nesterov,
259            stable_normalization: self.stable_normalization,
260            matrix_layout: self.matrix_layout,
261        })
262    }
263
264    /// Fallible model optimizer initialization. Only pass matrix-only modules;
265    /// use MuonAdamWConfig for a complete model with biases/embeddings.
266    pub fn try_init<B: AutodiffBackend, M: AutodiffModule<B>>(
267        &self,
268    ) -> Result<OptimizerAdaptor<Muon<B::InnerBackend>, M, B>, MuonError> {
269        Ok(OptimizerAdaptor::from(self.try_build()?))
270    }
271
272    /// Initialize Muon optimizer.
273    ///
274    /// # Returns
275    ///
276    /// Returns an optimizer adaptor that can be used to optimize a module.
277    ///
278    /// # Example
279    ///
280    /// ```ignore
281    /// use ruda_optim::{MuonConfig, AdjustLrFn, decay::WeightDecayConfig};
282    ///
283    /// // Basic configuration with default (Original) LR adjustment
284    /// let optimizer = MuonConfig::new()
285    ///     .with_weight_decay(Some(WeightDecayConfig::new(0.01)))
286    ///     .init();
287    ///
288    /// // With AdamW-compatible settings using MatchRmsAdamW
289    /// let optimizer = MuonConfig::new()
290    ///     .with_adjust_lr_fn(AdjustLrFn::MatchRmsAdamW)
291    ///     .with_weight_decay(Some(WeightDecayConfig::new(0.1)))
292    ///     .init();
293    ///
294    /// // Custom momentum and NS settings
295    /// let optimizer = MuonConfig::new()
296    ///     .with_momentum(MomentumConfig {
297    ///         momentum: 0.9,
298    ///         dampening: 0.1,
299    ///         nesterov: false,
300    ///     })
301    ///     .with_ns_steps(7)
302    ///     .init();
303    /// ```
304    pub fn init<B: AutodiffBackend, M: AutodiffModule<B>>(
305        &self,
306    ) -> OptimizerAdaptor<Muon<B::InnerBackend>, M, B> {
307        OptimizerAdaptor::from(self.build())
308    }
309}
310
311/// Parameters for Newton-Schulz orthogonalization.
312#[derive(Clone, Copy)]
313struct NewtonSchulzParams {
314    a: f32,
315    b: f32,
316    c: f32,
317    steps: usize,
318}
319
320impl NewtonSchulzParams {
321    fn new(coefficients: (f32, f32, f32), steps: usize) -> Self {
322        Self {
323            a: coefficients.0,
324            b: coefficients.1,
325            c: coefficients.2,
326            steps,
327        }
328    }
329}
330
331/// Muon optimizer.
332///
333/// Muon internally runs standard SGD-momentum, and then performs an orthogonalization
334/// post-processing step using a finite quintic Newton-Schulz iteration. This is an
335/// approximate spectral transformation, NOT an exact polar decomposition or a
336/// guarantee that U*U^T is identity. This implementation computes in the tensor's
337/// dtype; it does not silently cast to BF16 or create FP32 master parameters.
338///
339/// # Important Notes
340///
341/// 1. **Only for nonempty 2D parameters**: Muon is designed for hidden weight matrices. Use AdamW
342///    or SGD for biases, embeddings, and layer norms.
343///
344/// 2. **Learning rate adjustment**: Muon automatically adjusts the learning rate based
345///    on parameter shape. See [`AdjustLrFn`] for details.
346///
347/// 3. **Weight decay timing**: Unlike typical optimizers, Muon applies weight decay
348///    AFTER orthogonalization but uses the original (unadjusted) learning rate for it.
349#[derive(Clone)]
350pub struct Muon<B: Backend> {
351    momentum: Momentum<B>,
352    ns_params: NewtonSchulzParams,
353    weight_decay_penalty: Option<f32>,
354    epsilon: f32,
355    adjust_lr_fn: AdjustLrFn,
356    momentum_mode: MuonMomentumMode,
357    momentum_beta: f64,
358    nesterov: bool,
359    stable_normalization: bool,
360    matrix_layout: MuonMatrixLayout,
361}
362
363impl<B: Backend> Muon<B> {
364    /// Check tensor metadata before submitting kernels. Does not scan values for
365    /// NaN/Inf. Unscale and validate gradients before calling (on every rank).
366    pub fn validate_step<const D: usize>(
367        &self, lr: LearningRate, tensor: &Tensor<B, D>, grad: &Tensor<B, D>,
368        state: Option<&MuonState<B, D>>,
369    ) -> Result<(), MuonError> {
370        if D != 2 { return Err(MuonError::ExpectedMatrix { rank: D }); }
371        if !lr.is_finite() || lr < 0.0 {
372            return Err(MuonError::InvalidConfig("learning rate must be finite and nonnegative"));
373        }
374        let shape = tensor.shape();
375        if shape.iter().any(|dim| *dim == 0) { return Err(MuonError::EmptyMatrix); }
376        if shape != grad.shape() { return Err(MuonError::ShapeMismatch("gradient")); }
377        if tensor.dtype() != grad.dtype() { return Err(MuonError::DTypeMismatch("gradient")); }
378        if tensor.device() != grad.device() { return Err(MuonError::DeviceMismatch("gradient")); }
379        if self.stable_normalization && tensor.dtype() != DType::F32 {
380            return Err(MuonError::InvalidConfig("stable normalization requires FP32 tensors"));
381        }
382        if let Some(state) = state {
383            let buffer = state.momentum.velocity();
384            if shape != buffer.shape() { return Err(MuonError::ShapeMismatch("momentum")); }
385            if tensor.dtype() != buffer.dtype() { return Err(MuonError::DTypeMismatch("momentum")); }
386            if tensor.device() != buffer.device() { return Err(MuonError::DeviceMismatch("momentum")); }
387        }
388        let adjusted = self.adjust_lr(lr, &shape);
389        let decay = lr * self.weight_decay_penalty.unwrap_or(0.0) as f64;
390        if !adjusted.is_finite() || !decay.is_finite()
391            || (tensor.dtype() == DType::F32 && (!(adjusted as f32).is_finite() || !(decay as f32).is_finite())) {
392            return Err(MuonError::InvalidConfig("effective learning rate/decay overflows"));
393        }
394        Ok(())
395    }
396
397    /// Fallible metadata-checked update. Output and state may still be executing
398    /// asynchronously. An accepted submission is not proof of device completion.
399    ///
400    /// Uses a complete 2D matrix, never a flattened mixed-parameter buffer or
401    /// an arbitrary shard. Missing gradients must be skipped by the caller.
402    pub fn try_step<const D: usize>(
403        &self, lr: LearningRate, tensor: Tensor<B, D>, grad: Tensor<B, D>,
404        state: Option<MuonState<B, D>>,
405    ) -> Result<(Tensor<B, D>, Option<MuonState<B, D>>), MuonError> {
406        self.validate_step(lr, &tensor, &grad, state.as_ref())?;
407        let (update, momentum) = match self.momentum_mode {
408            MuonMomentumMode::Sgd => self.momentum.transform(grad, state.map(|s| s.momentum)),
409            MuonMomentumMode::Ema => {
410                let beta = self.momentum_beta;
411                let buffer = match state {
412                    Some(s) => s.momentum.velocity().clone().mul_scalar(beta)
413                        .add(grad.clone().mul_scalar(1.0 - beta)),
414                    None => grad.clone().mul_scalar(1.0 - beta),
415                };
416                let update = if self.nesterov {
417                    grad.mul_scalar(1.0 - beta).add(buffer.clone().mul_scalar(beta))
418                } else { buffer.clone() };
419                (update, MomentumState::new(buffer))
420            }
421        };
422        let update = self.zeropower_via_newtonschulz(update);
423        let adjusted_lr = self.adjust_lr(lr, &tensor.shape());
424        let tensor = match self.weight_decay_penalty {
425            Some(penalty) => tensor.mul_scalar(1.0 - lr * penalty as f64),
426            None => tensor,
427        };
428        Ok((tensor - update.mul_scalar(adjusted_lr), Some(MuonState::new(momentum))))
429    }
430
431    /// Adjust learning rate based on parameter shape.
432    ///
433    /// # Arguments
434    ///
435    /// * `lr` - Base learning rate
436    /// * `shape` - Parameter shape (uses first two dimensions)
437    ///
438    /// # Returns
439    ///
440    /// Adjusted learning rate
441    ///
442    /// ```ignore
443    /// // For a [1024, 512] weight matrix with lr=0.01:
444    /// // Original: 0.01 * sqrt(1024/512) = 0.01 * 1.414 = 0.01414
445    /// // MatchRmsAdamW: 0.01 * 0.2 * sqrt(1024) = 0.01 * 0.2 * 32 = 0.064
446    /// ```
447    fn adjust_lr(&self, lr: LearningRate, shape: &[usize]) -> LearningRate {
448        let ratio = match self.matrix_layout {
449            MuonMatrixLayout::InputOutput if shape.len() == 2 =>
450                self.adjust_lr_fn.adjustment_ratio(&[shape[1], shape[0]]),
451            _ => self.adjust_lr_fn.adjustment_ratio(shape),
452        };
453        lr * ratio
454    }
455
456    /// Perform Newton-Schulz orthogonalization on a gradient tensor.
457    ///
458    /// This computes the zeroth power (orthogonalization) of the input matrix G
459    /// using a quintic Newton-Schulz iteration.
460    ///
461    /// # Algorithm
462    ///
463    /// 1. Transpose if tall matrix (A > B)
464    /// 2. Normalize: X = X / ||X||
465    /// 3. For k steps:
466    ///    - A = X @ X^T
467    ///    - B = b*A + c*A^2
468    ///    - X = a*X + B@X
469    /// 4. Transpose back if needed
470    ///
471    /// # References
472    ///
473    /// - Original: https://github.com/KellerJordan/Muon/blob/master/muon.py
474    /// - PyTorch: https://github.com/pytorch/pytorch/blob/main/torch/optim/muon.py
475    fn zeropower_via_newtonschulz<const D: usize>(&self, g: Tensor<B, D>) -> Tensor<B, D> {
476        let shape = g.shape();
477        let dim_m2 = shape[D - 2];
478        let dim_m1 = shape[D - 1];
479
480        // Step 1: Transpose if tall matrix (more rows than columns)
481        let (mut x, needs_transpose) = if dim_m2 > dim_m1 {
482            (g.swap_dims(D - 2, D - 1), true)
483        } else {
484            (g, false)
485        };
486
487        // Step 2: Normalize by Frobenius norm
488        // X = X / (||X|| + epsilon)
489        if self.stable_normalization {
490            // Equivalent in real arithmetic to x / max(||x||_F, eps), without
491            // forming x*x at its original scale. A zero matrix stays zero.
492            // The caller validates FP32; MIN_POSITIVE is not representable in FP16.
493            let scale = x.clone().abs().max().clamp_min(f32::MIN_POSITIVE);
494            let scaled = x.div(scale.clone().unsqueeze());
495            let floor = scale.recip().mul_scalar(self.epsilon);
496            let norm = scaled.clone().square().sum().sqrt().max_pair(floor);
497            x = scaled.div(norm.unsqueeze());
498        } else {
499            // Preserve the existing RUDA normalization and dtype/rounding path.
500            let norm = x.clone().powf_scalar(2.0).sum().sqrt()
501                .clamp_min(self.epsilon).unsqueeze();
502            x = x.div(norm);
503        }
504
505        // Step 3: Newton-Schulz iteration
506        // This is the quintic iteration with coefficients (a, b, c)
507        let NewtonSchulzParams { a, b, c, steps } = self.ns_params;
508
509        for _ in 0..steps {
510            // A = X @ X^T
511            let x_t = x.clone().swap_dims(D - 2, D - 1);
512            let a_matrix = x.clone().matmul(x_t);
513
514            // B = b*A + c*A@A
515            let a_squared = a_matrix.clone().matmul(a_matrix.clone());
516            let b_matrix = a_matrix.mul_scalar(b).add(a_squared.mul_scalar(c));
517
518            // X = a*X + B@X
519            x = x.clone().mul_scalar(a).add(b_matrix.matmul(x.clone()));
520        }
521
522        // Step 4: Restore transpose if it was a tall matrix
523        if needs_transpose {
524            x = x.swap_dims(D - 2, D - 1);
525        }
526
527        x
528    }
529}
530
531/// Muon state.
532#[derive(Record, Clone, new)]
533pub struct MuonState<B: Backend, const D: usize> {
534    /// Current momentum state
535    pub momentum: MomentumState<B, D>,
536}
537
538impl<B: Backend> SimpleOptimizer<B> for Muon<B> {
539    type State<const D: usize> = MuonState<B, D>;
540
541    /// Perform a single Muon optimization step.
542    ///
543    /// # Algorithm
544    ///
545    /// 1. Apply momentum to gradient
546    /// 2. Orthogonalize update via Newton-Schulz
547    /// 3. Adjust learning rate based on parameter shape
548    /// 4. Apply weight decay (using original lr)
549    /// 5. Update parameter (using adjusted lr)
550    ///
551    /// # Notes
552    ///
553    /// Unlike typical optimizers, the weight decay and parameter update use
554    /// different learning rates:
555    /// - Weight decay uses the original `lr`
556    /// - Parameter update uses the shape-adjusted `lr`
557    ///
558    /// # Panics
559    /// This function will panic if the input tensors are not 2D.
560    fn step<const D: usize>(
561        &self,
562        lr: LearningRate,
563        tensor: Tensor<B, D>,
564        grad: Tensor<B, D>,
565        state: Option<Self::State<D>>,
566    ) -> (Tensor<B, D>, Option<Self::State<D>>) {
567        self.try_step(lr, tensor, grad, state)
568            .unwrap_or_else(|error| panic!("{error}"))
569    }
570
571    fn to_device<const D: usize>(mut state: Self::State<D>, device: &Device<B>) -> Self::State<D> {
572        state.momentum = state.momentum.to_device(device);
573        state
574    }
575}
576
577#[cfg(test)]
578mod tests;