Skip to content

Commit cc21885

Browse files
committed
remove muon scale factor since moso already has the adamW update
Signed-off-by: mikail <mkhona@nvidia.com>
1 parent fe9a22a commit cc21885

2 files changed

Lines changed: 0 additions & 12 deletions

File tree

emerging_optimizers/soap/moso.py

Lines changed: 0 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -24,7 +24,6 @@
2424

2525
from emerging_optimizers import mixin as opt_mixin
2626
from emerging_optimizers import registry, utils
27-
from emerging_optimizers.orthogonalized_optimizers.muon import MuonScaleT, get_muon_scale_factor
2827
from emerging_optimizers.scalar_optimizers import update_functions
2928
from emerging_optimizers.soap import soap_utils
3029
from 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

tests/test_moso.py

Lines changed: 0 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,6 @@
1717
from absl.testing import absltest, parameterized
1818

1919
from emerging_optimizers import registry
20-
from emerging_optimizers.orthogonalized_optimizers.muon import get_muon_scale_factor
2120
from emerging_optimizers.soap import MOSO
2221

2322

@@ -106,7 +105,6 @@ def test_no_ema_is_close_to_one_sided_adam_in_eigenbasis(self, shape: tuple[int,
106105
shampoo_beta=0.0,
107106
eps=1e-12,
108107
weight_decay=0.0,
109-
scale_mode="spectral",
110108
)
111109

112110
optimizer.step()
@@ -122,7 +120,6 @@ def test_no_ema_is_close_to_one_sided_adam_in_eigenbasis(self, shape: tuple[int,
122120
adam_projected = projected / (projected.abs() + optimizer.param_groups[0]["eps"])
123121
expected_update = adam_projected @ eigenbasis.T
124122

125-
expected_update = expected_update * get_muon_scale_factor(*shape, mode="spectral")
126123
applied_update = -param.detach() / lr
127124
torch.testing.assert_close(
128125
applied_update,

0 commit comments

Comments
 (0)