pub struct MuonAdamW<M: AutodiffModule<B>, B: AutodiffBackend> { /* private fields */ }Expand description
Mixed optimizer with fixed, explicit parameter identities.
Missing gradients skip both momentum and decay for that parameter.
A tied parameter is updated once, by its ParamId.
Implementations§
Source§impl<M: AutodiffModule<B>, B: AutodiffBackend> MuonAdamW<M, B>
impl<M: AutodiffModule<B>, B: AutodiffBackend> MuonAdamW<M, B>
Sourcepub fn muon_parameter_count(&self) -> usize
pub fn muon_parameter_count(&self) -> usize
Number of distinct parameters routed to Muon, not number of matrix elements.
Sourcepub fn try_step_with_lrs(
&mut self,
muon_lr: LearningRate,
adamw_lr: LearningRate,
module: M,
grads: GradientsParams,
) -> Result<M, MuonError>
pub fn try_step_with_lrs( &mut self, muon_lr: LearningRate, adamw_lr: LearningRate, module: M, grads: GradientsParams, ) -> Result<M, MuonError>
Update with independent learning rates. Metadata for both groups is checked before either group submits an update. Device runtime errors are still asynchronous and this method is NOT a two-phase device transaction.
Sourcepub fn try_step_or_skip(
&mut self,
muon_lr: LearningRate,
adamw_lr: LearningRate,
module: M,
grads: GradientsParams,
skip_update: bool,
) -> Result<M, MuonError>
pub fn try_step_or_skip( &mut self, muon_lr: LearningRate, adamw_lr: LearningRate, module: M, grads: GradientsParams, skip_update: bool, ) -> Result<M, MuonError>
An explicit all-group skip makes no optimizer update and changes no state. The caller must decide skip consistently across all replicas and must unscale/check finite gradients before a non-skipped update.
Sourcepub fn try_load_record(
self,
record: MuonAdamWRecord<B>,
) -> Result<Self, MuonError>
pub fn try_load_record( self, record: MuonAdamWRecord<B>, ) -> Result<Self, MuonError>
Restore only compatible grouping/configuration. Does not silently move a momentum buffer between SGD and EMA conventions or between parameter roles.