Repository navigation
Stage Muon's momentum once per partition in ZeRO-1/2 - #8633
Merged
pengdurice merged 2 commits intoSep 24, 2026
Merged
Conversation
_muon_staging_momentum copies the committed momentum into the staging buffer, and it was called inside the per-parameter loop of get_flat_partition and _get_flat_partition_unpadded. Each call wiped what the parameters before it had just written, so only the last Muon matrix of each partition kept its momentum. Stage once, before the loop. The step-0 check in test_a_discarded_step_leaves_every_momentum_where_it_was compared against a momentum buffer that does not exist before the first step, so it held whatever happened. It now requires both momenta to be non-zero. TestMuonMatchesZeroStage0 compares ZeRO-1/2 against ZeRO-0, which applies Muon per parameter in the optimizer. Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
alanhuangyoo
requested review from
loadams,
tjruwase and
tohtana
as code owners
September 23, 2026 04:40
pengdurice
enabled auto-merge
September 24, 2026 16:50
pengdurice
approved these changes
Sep 24, 2026
Contributor
Author
|
@pengdurice Thanks for setting auto-merge! It's still waiting on an approving review, in case that got missed. |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
#8435 (mine) broke ZeRO-1/2 Muon: every matrix except the last one in each partition loses its momentum.
_muon_staging_momentumcopies the committed momentum into the staging buffer, and it was called inside the per-parameter loop ofget_flat_partition/_get_flat_partition_unpadded. So each parameter wiped what the ones before it had just written, and only the last one's momentum got committed. Nothing errors and the loss still goes down. This stages once per partition, before the loop.Four matrices on 2×H20, fp32, the same batch on every rank, clipping off, 3 steps. ZeRO-0 applies Muon per parameter inside the optimizer, so it is the reference. Relative weight difference to ZeRO-0, per matrix:
The #8435 invariant still holds: staging still starts from the committed momentum, so a step the loss scaler throws away leaves every momentum where it was.
Tests:
TestMuonMatchesZeroStage0compares ZeRO-1/2 with ZeRO-0 after three steps. Fails on master (3.5e-2), passes here.test_a_discarded_step_leaves_every_momentum_where_it_waswas vacuous. The buffer does not exist before the first step, so it comparedNonewith the new norms. It now requires both momenta to be non-zero, and fails on master, where the first matrix stays at exactly 0.tests/unit/v1/ops/muonandtests/unit/runtime/zero/test_per_head_muon.py: 292 passed.Clipping is off in the comparison because the default clipping scales the Muon update by its own norm (#8439), and that norm is taken over different tensors at each stage.