2424
2525from emerging_optimizers import mixin as opt_mixin
2626from emerging_optimizers import registry , utils
27- from emerging_optimizers .orthogonalized_optimizers .muon import MuonScaleT , get_muon_scale_factor
2827from emerging_optimizers .scalar_optimizers import update_functions
2928from emerging_optimizers .soap import soap_utils
3029from emerging_optimizers .soap .soap import _clip_update_rms_in_place
@@ -61,8 +60,6 @@ class MOSO(opt_mixin.WeightDecayMixin, optim.Optimizer):
6160 shampoo_beta: EMA coefficient for the one-sided momentum covariance.
6261 eps: Inner Adam epsilon for numerical stability.
6362 weight_decay: Weight decay coefficient.
64- scale_mode: Muon update scale mode.
65- extra_scale_factor: Additional update scale factor.
6663 max_update_rms: Clip the update RMS to this value (0 means no clipping).
6764 """
6865
@@ -76,8 +73,6 @@ def __init__(
7673 eps : float = 1e-8 ,
7774 weight_decay : float = 0.01 ,
7875 * ,
79- scale_mode : MuonScaleT = "spectral" ,
80- extra_scale_factor : float = 1.0 ,
8176 max_update_rms : float = 0.0 ,
8277 ) -> None :
8378 self .nesterov = False
@@ -88,8 +83,6 @@ def __init__(
8883 self .use_eigh = False
8984 self .qr_fp32_matmul_prec : FP32MatmulPrecT = "highest"
9085 self .power_iter_steps = 1
91- self .scale_mode = scale_mode
92- self .extra_scale_factor = extra_scale_factor
9386 self .max_update_rms = max_update_rms
9487
9588 defaults = {
@@ -217,8 +210,6 @@ def step(self, closure: Callable[[], float] | None = None) -> float | None:
217210 left_preconditioned = left_preconditioned ,
218211 )
219212
220- scale_factor = get_muon_scale_factor (momentum .shape [0 ], momentum .shape [1 ], mode = self .scale_mode )
221- update = update * scale_factor * self .extra_scale_factor
222213 _clip_update_rms_in_place (update , self .max_update_rms )
223214 p .add_ (update , alpha = - group ["lr" ])
224215
0 commit comments