Skip to content

Commit 5bceeca

Browse files
committed
Addressed MR comments
Signed-off-by: mikail <mkhona@nvidia.com>
1 parent 7fe9f1e commit 5bceeca

2 files changed

Lines changed: 9 additions & 10 deletions

File tree

emerging_optimizers/soap/moso.py

Lines changed: 7 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -250,10 +250,9 @@ def _update_one_sided_momentum_factor(
250250
shampoo_beta: float,
251251
) -> None:
252252
"""Update the smaller-side covariance of the Muon momentum."""
253-
if momentum.shape[0] <= momentum.shape[1]:
254-
momentum_factor.lerp_(momentum @ momentum.T, 1 - shampoo_beta)
255-
else:
256-
momentum_factor.lerp_(momentum.T @ momentum, 1 - shampoo_beta)
253+
left_preconditioned = momentum.shape[0] <= momentum.shape[1]
254+
maybe_transposed_momentum = momentum if left_preconditioned else momentum.T
255+
momentum_factor.lerp_(maybe_transposed_momentum @ maybe_transposed_momentum.T, 1 - shampoo_beta)
257256

258257

259258
@torch.no_grad() # type: ignore[misc]
@@ -282,13 +281,13 @@ def _update_eigenbasis_and_adam_exp_avgs(
282281
)
283282

284283
if use_eigh:
285-
updated_eigenbasis = soap_utils.get_eigenbasis_eigh([momentum_factor])[0]
284+
(updated_eigenbasis,) = soap_utils.get_eigenbasis_eigh([momentum_factor])
286285
else:
287-
updated_eigenbasis = soap_utils.get_eigenbasis_qr(
286+
(updated_eigenbasis,) = soap_utils.get_eigenbasis_qr(
288287
[momentum_factor],
289288
[eigenbasis],
290289
power_iter_steps=power_iter_steps,
291-
)[0]
290+
)
292291

293292
exp_avg = _project_to_one_sided_eigenbasis(
294293
x=exp_avg,
@@ -307,7 +306,7 @@ def _sort_one_sided_eigenbasis_and_exp_avg_sq(
307306
) -> tuple[torch.Tensor, torch.Tensor]:
308307
"""Sort eigenbasis slots by approximate eigenvalue and permute Adam second moments."""
309308
approx_eigvals = utils.eig.conjugate(momentum_factor, eigenbasis, diag=True)
310-
sort_idx = torch.argsort(approx_eigvals, descending=True)
309+
sort_idx = torch.argsort(approx_eigvals, descending=True, stable=True)
311310
sorted_eigenbasis = eigenbasis[:, sort_idx]
312311
exp_avg_sq_dim = 0 if left_preconditioned else 1
313312
return sorted_eigenbasis, exp_avg_sq.index_select(exp_avg_sq_dim, sort_idx)

tests/test_moso.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -93,7 +93,7 @@ def test_accumulates_momentum_covariance_on_smaller_side(self, shape: tuple[int,
9393
{"shape": (4, 8)},
9494
{"shape": (8, 4)},
9595
)
96-
def test_no_ema_matches_one_sided_adam_in_eigenbasis(self, shape: tuple[int, int]) -> None:
96+
def test_no_ema_is_close_to_one_sided_adam_in_eigenbasis(self, shape: tuple[int, int]) -> None:
9797
torch.manual_seed(7)
9898
grad = torch.randn(shape, device=FLAGS.device)
9999
param = torch.zeros(shape, requires_grad=True, device=FLAGS.device)
@@ -131,7 +131,7 @@ def test_no_ema_matches_one_sided_adam_in_eigenbasis(self, shape: tuple[int, int
131131
torch.testing.assert_close(
132132
applied_update,
133133
expected_update,
134-
atol=1e-4,
134+
atol=0.0,
135135
rtol=1e-4,
136136
msg=lambda msg: f"MOSO no-EMA update did not match projected Adam update for shape {shape}:\n{msg}",
137137
)

0 commit comments

Comments
 (0)