@@ -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 )
0 commit comments