1use 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#[derive(Config, Debug)]
29pub struct MuonAdamWConfig {
30 #[config(default = "MuonConfig::new()")]
32 muon: MuonConfig,
33 #[config(default = "AdamWConfig::new()")]
35 adamw: AdamWConfig,
36 #[config(default = 0.015)]
38 adamw_lr_ratio: f64,
39}
40
41impl MuonAdamWConfig {
42 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#[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#[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 pub fn muon_state_count(&self) -> usize { self.muon.len() }
93 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 pub fn muon_parameter_count(&self) -> usize { self.selected.len() }
120
121 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 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 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 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(¶m.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(¶m.id);
285 if selected && !trainable { self.error = Some(MuonError::FrozenParameter(param.id.val())); return; }
286 let inner = tensor.inner();
287 if selected {
288 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(¶m.id) { self.muon.register(param.id, grad); }
316 else { self.adamw.register(param.id, grad); }
317 }
318 }
319}
320
321#[cfg(test)]
322mod tests;