use ruda_model::{module::AutodiffModule, record::Record};
use ruda_model::config::Config;
use ruda_model::tensor::{Tensor, backend::AutodiffBackend};
use ruda_model::tensor::{backend::Backend, ops::Device};
use serde::{Deserialize, Serialize};
use super::{
SimpleOptimizer,
adaptor::OptimizerAdaptor,
decay::WeightDecayConfig,
momentum::{Momentum, MomentumConfig, MomentumState},
};
use crate::LearningRate;
use ruda_model::tensor::DType;
mod error;
pub use error::MuonError;
mod grouped;
pub use grouped::{MuonAdamW, MuonAdamWConfig, MuonAdamWRecord};
#[derive(Clone, Default, Debug, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum MuonMomentumMode {
#[default]
Sgd,
Ema,
}
#[cfg(not(feature = "std"))]
#[allow(unused_imports)]
use num_traits::Float as _;
#[derive(Clone, Default, Debug, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum MuonMatrixLayout {
#[default]
AsStored,
InputOutput,
}
#[derive(Clone, Default, Debug, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum AdjustLrFn {
#[default]
Original,
MatchRmsAdamW,
}
impl AdjustLrFn {
fn adjustment_ratio(&self, shape: &[usize]) -> f64 {
if shape.len() < 2 {
return 1.0;
}
let a = shape[0] as f64;
let b = shape[1] as f64;
match self {
Self::Original => {
let ratio = a / b;
ratio.max(1.0).sqrt()
}
Self::MatchRmsAdamW => {
0.2 * a.max(b).sqrt()
}
}
}
}
#[derive(Config, Debug)]
pub struct MuonConfig {
weight_decay: Option<WeightDecayConfig>,
#[config(default = "MomentumConfig { momentum: 0.95, dampening: 0.0, nesterov: true }")]
momentum: MomentumConfig,
#[config(default = "(3.4445, -4.775, 2.0315)")]
ns_coefficients: (f32, f32, f32),
#[config(default = 1e-7)]
epsilon: f32,
#[config(default = 5)]
ns_steps: usize,
#[config(default = "AdjustLrFn::Original")]
adjust_lr_fn: AdjustLrFn,
#[config(default = "MuonMomentumMode::Sgd")]
momentum_mode: MuonMomentumMode,
#[config(default = false)]
stable_normalization: bool,
#[config(default = "MuonMatrixLayout::AsStored")]
matrix_layout: MuonMatrixLayout,
}
impl MuonConfig {
pub fn build<B: Backend>(&self) -> Muon<B> {
self.try_build().unwrap_or_else(|error| panic!("{error}"))
}
pub fn validate(&self) -> Result<(), MuonError> {
let beta = self.momentum.momentum;
let dampening = self.momentum.dampening;
if !beta.is_finite() || !(0.0..1.0).contains(&beta) {
return Err(MuonError::InvalidConfig("momentum must be finite in [0, 1)"));
}
if !dampening.is_finite() || !(0.0..=1.0).contains(&dampening) {
return Err(MuonError::InvalidConfig("dampening must be finite in [0, 1]"));
}
if self.momentum.nesterov && (beta == 0.0 || dampening != 0.0) {
return Err(MuonError::InvalidConfig("Nesterov requires positive momentum and zero dampening"));
}
if self.momentum_mode == MuonMomentumMode::Ema && dampening != 0.0 {
return Err(MuonError::InvalidConfig("EMA momentum requires zero dampening"));
}
if !self.epsilon.is_finite() || self.epsilon <= 0.0 {
return Err(MuonError::InvalidConfig("epsilon must be finite and positive"));
}
if !(1..100).contains(&self.ns_steps) {
return Err(MuonError::InvalidConfig("Newton-Schulz steps must be in 1..100"));
}
let (a, b, c) = self.ns_coefficients;
if !a.is_finite() || !b.is_finite() || !c.is_finite() {
return Err(MuonError::InvalidConfig("Newton-Schulz coefficients must be finite"));
}
if let Some(decay) = &self.weight_decay {
if !decay.penalty.is_finite() || decay.penalty < 0.0 {
return Err(MuonError::InvalidConfig("weight decay must be finite and nonnegative"));
}
}
Ok(())
}
pub fn try_build<B: Backend>(&self) -> Result<Muon<B>, MuonError> {
self.validate()?;
let momentum = Momentum::new(&self.momentum);
let weight_decay_penalty = self.weight_decay.as_ref().map(|wd| wd.penalty);
Ok(Muon {
momentum,
ns_params: NewtonSchulzParams::new(self.ns_coefficients, self.ns_steps),
weight_decay_penalty,
epsilon: self.epsilon,
adjust_lr_fn: self.adjust_lr_fn,
momentum_mode: self.momentum_mode,
momentum_beta: self.momentum.momentum,
nesterov: self.momentum.nesterov,
stable_normalization: self.stable_normalization,
matrix_layout: self.matrix_layout,
})
}
pub fn try_init<B: AutodiffBackend, M: AutodiffModule<B>>(
&self,
) -> Result<OptimizerAdaptor<Muon<B::InnerBackend>, M, B>, MuonError> {
Ok(OptimizerAdaptor::from(self.try_build()?))
}
pub fn init<B: AutodiffBackend, M: AutodiffModule<B>>(
&self,
) -> OptimizerAdaptor<Muon<B::InnerBackend>, M, B> {
OptimizerAdaptor::from(self.build())
}
}
#[derive(Clone, Copy)]
struct NewtonSchulzParams {
a: f32,
b: f32,
c: f32,
steps: usize,
}
impl NewtonSchulzParams {
fn new(coefficients: (f32, f32, f32), steps: usize) -> Self {
Self {
a: coefficients.0,
b: coefficients.1,
c: coefficients.2,
steps,
}
}
}
#[derive(Clone)]
pub struct Muon<B: Backend> {
momentum: Momentum<B>,
ns_params: NewtonSchulzParams,
weight_decay_penalty: Option<f32>,
epsilon: f32,
adjust_lr_fn: AdjustLrFn,
momentum_mode: MuonMomentumMode,
momentum_beta: f64,
nesterov: bool,
stable_normalization: bool,
matrix_layout: MuonMatrixLayout,
}
impl<B: Backend> Muon<B> {
pub fn validate_step<const D: usize>(
&self, lr: LearningRate, tensor: &Tensor<B, D>, grad: &Tensor<B, D>,
state: Option<&MuonState<B, D>>,
) -> Result<(), MuonError> {
if D != 2 { return Err(MuonError::ExpectedMatrix { rank: D }); }
if !lr.is_finite() || lr < 0.0 {
return Err(MuonError::InvalidConfig("learning rate must be finite and nonnegative"));
}
let shape = tensor.shape();
if shape.iter().any(|dim| *dim == 0) { return Err(MuonError::EmptyMatrix); }
if shape != grad.shape() { return Err(MuonError::ShapeMismatch("gradient")); }
if tensor.dtype() != grad.dtype() { return Err(MuonError::DTypeMismatch("gradient")); }
if tensor.device() != grad.device() { return Err(MuonError::DeviceMismatch("gradient")); }
if self.stable_normalization && tensor.dtype() != DType::F32 {
return Err(MuonError::InvalidConfig("stable normalization requires FP32 tensors"));
}
if let Some(state) = state {
let buffer = state.momentum.velocity();
if shape != buffer.shape() { return Err(MuonError::ShapeMismatch("momentum")); }
if tensor.dtype() != buffer.dtype() { return Err(MuonError::DTypeMismatch("momentum")); }
if tensor.device() != buffer.device() { return Err(MuonError::DeviceMismatch("momentum")); }
}
let adjusted = self.adjust_lr(lr, &shape);
let decay = lr * self.weight_decay_penalty.unwrap_or(0.0) as f64;
if !adjusted.is_finite() || !decay.is_finite()
|| (tensor.dtype() == DType::F32 && (!(adjusted as f32).is_finite() || !(decay as f32).is_finite())) {
return Err(MuonError::InvalidConfig("effective learning rate/decay overflows"));
}
Ok(())
}
pub fn try_step<const D: usize>(
&self, lr: LearningRate, tensor: Tensor<B, D>, grad: Tensor<B, D>,
state: Option<MuonState<B, D>>,
) -> Result<(Tensor<B, D>, Option<MuonState<B, D>>), MuonError> {
self.validate_step(lr, &tensor, &grad, state.as_ref())?;
let (update, momentum) = match self.momentum_mode {
MuonMomentumMode::Sgd => self.momentum.transform(grad, state.map(|s| s.momentum)),
MuonMomentumMode::Ema => {
let beta = self.momentum_beta;
let buffer = match state {
Some(s) => s.momentum.velocity().clone().mul_scalar(beta)
.add(grad.clone().mul_scalar(1.0 - beta)),
None => grad.clone().mul_scalar(1.0 - beta),
};
let update = if self.nesterov {
grad.mul_scalar(1.0 - beta).add(buffer.clone().mul_scalar(beta))
} else { buffer.clone() };
(update, MomentumState::new(buffer))
}
};
let update = self.zeropower_via_newtonschulz(update);
let adjusted_lr = self.adjust_lr(lr, &tensor.shape());
let tensor = match self.weight_decay_penalty {
Some(penalty) => tensor.mul_scalar(1.0 - lr * penalty as f64),
None => tensor,
};
Ok((tensor - update.mul_scalar(adjusted_lr), Some(MuonState::new(momentum))))
}
fn adjust_lr(&self, lr: LearningRate, shape: &[usize]) -> LearningRate {
let ratio = match self.matrix_layout {
MuonMatrixLayout::InputOutput if shape.len() == 2 =>
self.adjust_lr_fn.adjustment_ratio(&[shape[1], shape[0]]),
_ => self.adjust_lr_fn.adjustment_ratio(shape),
};
lr * ratio
}
fn zeropower_via_newtonschulz<const D: usize>(&self, g: Tensor<B, D>) -> Tensor<B, D> {
let shape = g.shape();
let dim_m2 = shape[D - 2];
let dim_m1 = shape[D - 1];
let (mut x, needs_transpose) = if dim_m2 > dim_m1 {
(g.swap_dims(D - 2, D - 1), true)
} else {
(g, false)
};
if self.stable_normalization {
let scale = x.clone().abs().max().clamp_min(f32::MIN_POSITIVE);
let scaled = x.div(scale.clone().unsqueeze());
let floor = scale.recip().mul_scalar(self.epsilon);
let norm = scaled.clone().square().sum().sqrt().max_pair(floor);
x = scaled.div(norm.unsqueeze());
} else {
let norm = x.clone().powf_scalar(2.0).sum().sqrt()
.clamp_min(self.epsilon).unsqueeze();
x = x.div(norm);
}
let NewtonSchulzParams { a, b, c, steps } = self.ns_params;
for _ in 0..steps {
let x_t = x.clone().swap_dims(D - 2, D - 1);
let a_matrix = x.clone().matmul(x_t);
let a_squared = a_matrix.clone().matmul(a_matrix.clone());
let b_matrix = a_matrix.mul_scalar(b).add(a_squared.mul_scalar(c));
x = x.clone().mul_scalar(a).add(b_matrix.matmul(x.clone()));
}
if needs_transpose {
x = x.swap_dims(D - 2, D - 1);
}
x
}
}
#[derive(Record, Clone, new)]
pub struct MuonState<B: Backend, const D: usize> {
pub momentum: MomentumState<B, D>,
}
impl<B: Backend> SimpleOptimizer<B> for Muon<B> {
type State<const D: usize> = MuonState<B, D>;
fn step<const D: usize>(
&self,
lr: LearningRate,
tensor: Tensor<B, D>,
grad: Tensor<B, D>,
state: Option<Self::State<D>>,
) -> (Tensor<B, D>, Option<Self::State<D>>) {
self.try_step(lr, tensor, grad, state)
.unwrap_or_else(|error| panic!("{error}"))
}
fn to_device<const D: usize>(mut state: Self::State<D>, device: &Device<B>) -> Self::State<D> {
state.momentum = state.momentum.to_device(device);
state
}
}
#[cfg(test)]
mod tests;