Skip to main content

ruda_optim/optim/muon/
error.rs

1// SPDX-License-Identifier: Apache-2.0
2use core::fmt;
3
4/// Configuration, metadata or explicit parameter-group error for Muon.
5/// Device execution faults remain backend errors; no host fallback is attempted.
6#[derive(Clone, Debug, PartialEq, Eq)]
7pub enum MuonError {
8    /// Invalid hyperparameter or unsupported numerical mode.
9    InvalidConfig(&'static str),
10    /// Muon operates on a complete matrix, not an arbitrary-rank tensor.
11    ExpectedMatrix {
12        /// Actual tensor rank.
13        rank: usize,
14    },
15    /// A matrix dimension is zero.
16    EmptyMatrix,
17    /// Gradient or restored momentum has incompatible geometry.
18    ShapeMismatch(&'static str),
19    /// Gradient or restored momentum has a different element format.
20    DTypeMismatch(&'static str),
21    /// Gradient or restored momentum is on a different device.
22    DeviceMismatch(&'static str),
23    /// The Muon group was left empty.
24    EmptyMuonGroup,
25    /// A selected id was supplied twice.
26    DuplicateParameter(u64),
27    /// A selected id is absent from the module.
28    UnknownParameter(u64),
29    /// A selected Muon parameter is frozen.
30    FrozenParameter(u64),
31    /// A module's ids/shapes changed, or tied aliases disagree.
32    ModelChanged,
33    /// A gradient refers to an unknown or frozen parameter.
34    UnusedGradients,
35    /// A record belongs to another grouping, geometry, configuration or schema.
36    IncompatibleRecord,
37    /// This API is deliberately limited to full, synchronized local matrices.
38    UnsupportedDistributed,
39}
40
41impl fmt::Display for MuonError {
42    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
43        match self {
44            Self::InvalidConfig(message) => write!(f, "invalid Muon configuration: {message}"),
45            Self::ExpectedMatrix { rank } => write!(f, "Newton-Schulz iteration requires 2D tensors, got {rank}D"),
46            Self::EmptyMatrix => f.write_str("Muon requires a nonempty matrix"),
47            Self::ShapeMismatch(what) => write!(f, "Muon {what} shape does not match its parameter"),
48            Self::DTypeMismatch(what) => write!(f, "Muon {what} dtype does not match its parameter"),
49            Self::DeviceMismatch(what) => write!(f, "Muon {what} device does not match its parameter"),
50            Self::EmptyMuonGroup => f.write_str("select at least one hidden matrix explicitly for Muon"),
51            Self::DuplicateParameter(id) => write!(f, "duplicate Muon parameter id {id}"),
52            Self::UnknownParameter(id) => write!(f, "unknown Muon parameter id {id}"),
53            Self::FrozenParameter(id) => write!(f, "Muon parameter {id} does not require gradients"),
54            Self::ModelChanged => f.write_str("model parameter ids/shapes changed or tied aliases disagree"),
55            Self::UnusedGradients => f.write_str("gradients include an unknown or frozen parameter"),
56            Self::IncompatibleRecord => f.write_str("Muon/AdamW record configuration, grouping, geometry or version mismatch"),
57            Self::UnsupportedDistributed => f.write_str("Muon requires complete synchronized matrices; implicit sharded/step_multi updates are not supported"),
58        }
59    }
60}
61
62impl core::error::Error for MuonError {}