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;