Skip to main content

ruda_optim/optim/muon/
grouped.rs

1// SPDX-License-Identifier: Apache-2.0
2//! Explicit parameter routing. Reuses existing optimizers; never guesses roles
3//! from rank alone (embeddings and output heads are also matrices).
4use alloc::{format, string::String, vec::Vec};
5use core::marker::PhantomData;
6use hashbrown::{HashMap, HashSet};
7use ruda_model::{
8    config::Config,
9    module::{AutodiffModule, ModuleVisitor, Param, ParamId},
10    record::{PrecisionSettings, Record},
11    tensor::{Tensor, backend::AutodiffBackend},
12};
13use crate::{
14    AdamW, AdamWConfig, GradientsParams, LearningRate, MultiGradientsParams,
15    Optimizer, adaptor::OptimizerAdaptor, record::{AdaptorRecord, AdaptorRecordV1},
16};
17use super::{Muon, MuonConfig, MuonError};
18
19type Manifest = Vec<(u64, Vec<usize>, bool, String)>;
20type MuonRecords<B: AutodiffBackend> = HashMap<ParamId, AdaptorRecord<Muon<<B as AutodiffBackend>::InnerBackend>, B>>;
21type AdamRecords<B: AutodiffBackend> = HashMap<ParamId, AdaptorRecord<AdamW, B>>;
22
23/// Muon for explicitly selected hidden matrices and AdamW for the remainder.
24///
25/// `Optimizer::step(lr, ...)` uses `lr` for Muon and `lr * adamw_lr_ratio`
26/// for AdamW. `try_step_with_lrs` accepts independent rates instead.
27/// This is a high-level tensor optimizer, not the experimental fused AdamW API.
28#[derive(Config, Debug)]
29pub struct MuonAdamWConfig {
30    /// Muon settings, with legacy numerical defaults preserved.
31    #[config(default = "MuonConfig::new()")]
32    muon: MuonConfig,
33    /// Settings for all parameters not explicitly selected for Muon.
34    #[config(default = "AdamWConfig::new()")]
35    adamw: AdamWConfig,
36    /// AdamW/Muon learning-rate ratio. 0.015 maps 0.02 to 0.0003.
37    #[config(default = 0.015)]
38    adamw_lr_ratio: f64,
39}
40
41impl MuonAdamWConfig {
42    /// Validate configuration and resolve selected parameter IDs against a model.
43    /// Do not select embeddings, output heads, normalization gains or biases.
44    pub fn init<B: AutodiffBackend, M: AutodiffModule<B>>(
45        &self, module: &M, muon_parameters: &[ParamId],
46    ) -> Result<MuonAdamW<M, B>, MuonError> {
47        self.muon.validate()?;
48        self.adamw.validate_hyperparameters().map_err(MuonError::InvalidConfig)?;
49        valid_lr(self.adamw_lr_ratio)?;
50        if muon_parameters.is_empty() { return Err(MuonError::EmptyMuonGroup); }
51        let mut selected = HashSet::new();
52        for id in muon_parameters {
53            if !selected.insert(*id) { return Err(MuonError::DuplicateParameter(id.val())); }
54        }
55        let muon: OptimizerAdaptor<Muon<B::InnerBackend>, M, B> = self.muon.try_init()?;
56        let manifest = inspect_module(module, &selected, muon.optim(), None, 0.0)?;
57        Ok(MuonAdamW {
58            muon, adamw: self.adamw.init(), selected, manifest,
59            config_key: format!("muon-adamw-v1:{self:?}"),
60            adamw_lr_ratio: self.adamw_lr_ratio,
61        })
62    }
63}
64
65/// Mixed optimizer with fixed, explicit parameter identities.
66/// Missing gradients skip both momentum and decay for that parameter.
67/// A tied parameter is updated once, by its `ParamId`.
68#[derive(Clone)]
69pub struct MuonAdamW<M: AutodiffModule<B>, B: AutodiffBackend> {
70    muon: OptimizerAdaptor<Muon<B::InnerBackend>, M, B>,
71    adamw: OptimizerAdaptor<AdamW, M, B>,
72    selected: HashSet<ParamId>,
73    manifest: Manifest,
74    config_key: String,
75    adamw_lr_ratio: f64,
76}
77
78/// Versioned optimizer record. Save the model record with it so ParamIds survive.
79/// Configuration and routing are checked on load; lower-precision record settings
80/// can round momentum, so use full-precision records for continuation comparisons.
81#[derive(Clone)]
82pub struct MuonAdamWRecord<B: AutodiffBackend> {
83    version: u32,
84    config_key: String,
85    manifest: Manifest,
86    muon: MuonRecords<B>,
87    adamw: AdamRecords<B>,
88}
89
90impl<B: AutodiffBackend> MuonAdamWRecord<B> {
91    /// Number of hidden-matrix momentum states saved.
92    pub fn muon_state_count(&self) -> usize { self.muon.len() }
93    /// Number of auxiliary AdamW parameter states saved.
94    pub fn adamw_state_count(&self) -> usize { self.adamw.len() }
95}
96
97impl<B: AutodiffBackend> Record<B> for MuonAdamWRecord<B> {
98    type Item<S: PrecisionSettings> = (
99        u32, String, Manifest,
100        <MuonRecords<B> as Record<B>>::Item<S>,
101        <AdamRecords<B> as Record<B>>::Item<S>,
102    );
103    fn into_item<S: PrecisionSettings>(self) -> Self::Item<S> {
104        (self.version, self.config_key, self.manifest,
105         <MuonRecords<B> as Record<B>>::into_item::<S>(self.muon),
106         <AdamRecords<B> as Record<B>>::into_item::<S>(self.adamw))
107    }
108    fn from_item<S: PrecisionSettings>(item: Self::Item<S>, device: &B::Device) -> Self {
109        Self {
110            version: item.0, config_key: item.1, manifest: item.2,
111            muon: <MuonRecords<B> as Record<B>>::from_item::<S>(item.3, device),
112            adamw: <AdamRecords<B> as Record<B>>::from_item::<S>(item.4, device),
113        }
114    }
115}
116
117impl<M: AutodiffModule<B>, B: AutodiffBackend> MuonAdamW<M, B> {
118    /// Number of distinct parameters routed to Muon, not number of matrix elements.
119    pub fn muon_parameter_count(&self) -> usize { self.selected.len() }
120
121    /// Update with independent learning rates. Metadata for both groups is
122    /// checked before either group submits an update. Device runtime errors are
123    /// still asynchronous and this method is NOT a two-phase device transaction.
124    pub fn try_step_with_lrs(
125        &mut self, muon_lr: LearningRate, adamw_lr: LearningRate,
126        module: M, grads: GradientsParams,
127    ) -> Result<M, MuonError> {
128        valid_lr(muon_lr)?;
129        valid_lr(adamw_lr)?;
130        let manifest = inspect_module(&module, &self.selected, self.muon.optim(), Some(&grads), muon_lr)?;
131        if manifest != self.manifest { return Err(MuonError::ModelChanged); }
132        let mut split = SplitGradients::<B> {
133            selected: &self.selected, source: grads,
134            muon: GradientsParams::new(), adamw: GradientsParams::new(),
135            seen: HashSet::new(), backend: PhantomData,
136        };
137        module.visit(&mut split);
138        if !split.source.is_empty() { return Err(MuonError::UnusedGradients); }
139        let module = self.muon.step(muon_lr, module, split.muon);
140        Ok(self.adamw.step(adamw_lr, module, split.adamw))
141    }
142
143    /// An explicit all-group skip makes no optimizer update and changes no state.
144    /// The caller must decide skip consistently across all replicas and must
145    /// unscale/check finite gradients before a non-skipped update.
146    pub fn try_step_or_skip(
147        &mut self, muon_lr: LearningRate, adamw_lr: LearningRate,
148        module: M, grads: GradientsParams, skip_update: bool,
149    ) -> Result<M, MuonError> {
150        if skip_update { return Ok(module); }
151        self.try_step_with_lrs(muon_lr, adamw_lr, module, grads)
152    }
153
154    /// Restore only compatible grouping/configuration. Does not silently move a
155    /// momentum buffer between SGD and EMA conventions or between parameter roles.
156    pub fn try_load_record(mut self, record: MuonAdamWRecord<B>) -> Result<Self, MuonError> {
157        if record.version != 1 || record.config_key != self.config_key || record.manifest != self.manifest {
158            return Err(MuonError::IncompatibleRecord);
159        }
160        let known: HashSet<_> = self.manifest.iter().map(|(id, _, _, _)| ParamId::from(*id)).collect();
161        if record.muon.keys().any(|id| !self.selected.contains(id))
162            || record.adamw.keys().any(|id| !known.contains(id) || self.selected.contains(id)) {
163            return Err(MuonError::IncompatibleRecord);
164        }
165        for (id, state) in &record.muon {
166            let expected = self.manifest.iter().find(|entry| entry.0 == id.val())
167                .ok_or(MuonError::IncompatibleRecord)?;
168            match state {
169                AdaptorRecord::V1(AdaptorRecordV1::Rank2(state)) => {
170                    let v = state.momentum.velocity();
171                    if v.shape().to_vec() != expected.1 || format!("{:?}", v.dtype()) != expected.3 {
172                        return Err(MuonError::IncompatibleRecord);
173                    }
174                }
175                _ => return Err(MuonError::IncompatibleRecord),
176            }
177        }
178        for (id, state) in &record.adamw {
179            let expected = self.manifest.iter().find(|entry| entry.0 == id.val())
180                .ok_or(MuonError::IncompatibleRecord)?;
181            macro_rules! check_adam {
182                ($state:expr) => {{
183                    let m = &$state.momentum;
184                    if m.time == 0 || m.moment_1.shape().to_vec() != expected.1
185                        || m.moment_2.shape().to_vec() != expected.1
186                        || format!("{:?}", m.moment_1.dtype()) != expected.3
187                        || format!("{:?}", m.moment_2.dtype()) != expected.3
188                        || m.max_moment_2.as_ref().is_some_and(|v|
189                            v.shape().to_vec() != expected.1 || format!("{:?}", v.dtype()) != expected.3) {
190                        return Err(MuonError::IncompatibleRecord);
191                    }
192                }};
193            }
194            match state {
195                AdaptorRecord::V1(state) => match state {
196                    AdaptorRecordV1::Rank0(v) => check_adam!(v),
197                    AdaptorRecordV1::Rank1(v) => check_adam!(v),
198                    AdaptorRecordV1::Rank2(v) => check_adam!(v),
199                    AdaptorRecordV1::Rank3(v) => check_adam!(v),
200                    AdaptorRecordV1::Rank4(v) => check_adam!(v),
201                    AdaptorRecordV1::Rank5(v) => check_adam!(v),
202                    AdaptorRecordV1::Rank6(v) => check_adam!(v),
203                    AdaptorRecordV1::Rank7(v) => check_adam!(v),
204                    AdaptorRecordV1::Rank8(v) => check_adam!(v),
205                },
206            }
207        }
208        self.muon = self.muon.load_record(record.muon);
209        self.adamw = self.adamw.load_record(record.adamw);
210        Ok(self)
211    }
212}
213
214impl<M: AutodiffModule<B>, B: AutodiffBackend> Optimizer<M, B> for MuonAdamW<M, B> {
215    type Record = MuonAdamWRecord<B>;
216    fn step(&mut self, lr: LearningRate, module: M, grads: GradientsParams) -> M {
217        self.try_step_with_lrs(lr, lr * self.adamw_lr_ratio, module, grads)
218            .unwrap_or_else(|error| panic!("{error}"))
219    }
220    fn step_multi(&mut self, _lr: LearningRate, _module: M, _grads: MultiGradientsParams) -> M {
221        // Orthogonalization is nonlinear: orthogonalizing each shard separately
222        // is not equivalent to orthogonalizing the full matrix.
223        panic!("{}", MuonError::UnsupportedDistributed)
224    }
225    fn to_record(&self) -> Self::Record {
226        MuonAdamWRecord {
227            version: 1, config_key: self.config_key.clone(), manifest: self.manifest.clone(),
228            muon: self.muon.to_record(), adamw: self.adamw.to_record(),
229        }
230    }
231    fn load_record(self, record: Self::Record) -> Self {
232        self.try_load_record(record).unwrap_or_else(|error| panic!("{error}"))
233    }
234}
235
236fn valid_lr(lr: f64) -> Result<(), MuonError> {
237    if !lr.is_finite() || lr < 0.0 || !(lr as f32).is_finite() {
238        Err(MuonError::InvalidConfig("learning rates/ratios must be finite, nonnegative and FP32-representable"))
239    } else { Ok(()) }
240}
241
242struct Inspect<'a, B: AutodiffBackend> {
243    selected: &'a HashSet<ParamId>,
244    muon: &'a Muon<B::InnerBackend>,
245    grads: Option<&'a GradientsParams>,
246    lr: f64,
247    entries: HashMap<ParamId, (Vec<usize>, bool, String)>,
248    seen_grads: usize,
249    error: Option<MuonError>,
250}
251
252fn inspect_module<B: AutodiffBackend, M: AutodiffModule<B>>(
253    module: &M, selected: &HashSet<ParamId>, muon: &Muon<B::InnerBackend>,
254    grads: Option<&GradientsParams>, lr: f64,
255) -> Result<Manifest, MuonError> {
256    let mut inspect = Inspect { selected, muon, grads, lr, entries: HashMap::new(), seen_grads: 0, error: None };
257    module.visit(&mut inspect);
258    if let Some(error) = inspect.error { return Err(error); }
259    for id in selected {
260        if !inspect.entries.contains_key(id) { return Err(MuonError::UnknownParameter(id.val())); }
261    }
262    if grads.is_some_and(|g| g.len() != inspect.seen_grads) { return Err(MuonError::UnusedGradients); }
263    let mut manifest: Manifest = inspect.entries.into_iter()
264        .map(|(id, (shape, _, dtype))| (id.val(), shape, selected.contains(&id), dtype)).collect();
265    manifest.sort_by_key(|entry| entry.0);
266    Ok(manifest)
267}
268
269impl<B: AutodiffBackend> ModuleVisitor<B> for Inspect<'_, B> {
270    fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<B, D>>) {
271        if self.error.is_some() { return; }
272        let tensor = param.val();
273        let shape = tensor.shape().to_vec();
274        let trainable = param.is_require_grad();
275        let dtype = format!("{:?}", tensor.dtype());
276        if let Some(prior) = self.entries.get(&param.id) {
277            if prior != &(shape, trainable, dtype) { self.error = Some(MuonError::ModelChanged); }
278            return;
279        }
280        self.entries.insert(param.id, (shape, trainable, dtype));
281        if D > 8 { self.error = Some(MuonError::InvalidConfig("optimizer records support at most rank 8")); return; }
282        #[cfg(feature = "distributed")]
283        if tensor.is_distributed() { self.error = Some(MuonError::UnsupportedDistributed); return; }
284        let selected = self.selected.contains(&param.id);
285        if selected && !trainable { self.error = Some(MuonError::FrozenParameter(param.id.val())); return; }
286        let inner = tensor.inner();
287        if selected {
288            // Also validates rank/empty shape/stable-normalization dtype at init.
289            if let Err(error) = self.muon.validate_step(self.lr, &inner, &inner, None) {
290                self.error = Some(error); return;
291            }
292        }
293        if !trainable { return; }
294        if let Some(grad) = self.grads.and_then(|g| g.get::<B::InnerBackend, D>(param.id)) {
295            self.seen_grads += 1;
296            if inner.shape() != grad.shape() { self.error = Some(MuonError::ShapeMismatch("gradient")); return; }
297            if inner.dtype() != grad.dtype() { self.error = Some(MuonError::DTypeMismatch("gradient")); return; }
298            if inner.device() != grad.device() { self.error = Some(MuonError::DeviceMismatch("gradient")); }
299        }
300    }
301}
302
303struct SplitGradients<'a, B: AutodiffBackend> {
304    selected: &'a HashSet<ParamId>,
305    source: GradientsParams,
306    muon: GradientsParams,
307    adamw: GradientsParams,
308    seen: HashSet<ParamId>,
309    backend: PhantomData<B>,
310}
311impl<B: AutodiffBackend> ModuleVisitor<B> for SplitGradients<'_, B> {
312    fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<B, D>>) {
313        if !param.is_require_grad() || !self.seen.insert(param.id) { return; }
314        if let Some(grad) = self.source.remove::<B::InnerBackend, D>(param.id) {
315            if self.selected.contains(&param.id) { self.muon.register(param.id, grad); }
316            else { self.adamw.register(param.id, grad); }
317        }
318    }
319}
320
321#[cfg(test)]
322mod tests;