Skip to content

Stage Muon's momentum once per partition in ZeRO-1/2 - #8633

Merged
pengdurice merged 2 commits into
deepspeedai:masterfrom
alanhuangyoo:fix/muon-stage-momentum-once
Sep 24, 2026
Merged

pengdurice merged 2 commits into
deepspeedai:masterfrom
alanhuangyoo:fix/muon-stage-momentum-once

Conversation

@alanhuangyoo

Copy link
Copy Markdown
Contributor

#8435 (mine) broke ZeRO-1/2 Muon: every matrix except the last one in each partition loses its momentum.

_muon_staging_momentum copies the committed momentum into the staging buffer, and it was called inside the per-parameter loop of get_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:

master this PR
ZeRO-1 2.0e-3 to 3.4e-2 0
ZeRO-2 2.0e-3 to 3.4e-2 0
ZeRO-3 0 0

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:

  • New TestMuonMatchesZeroStage0 compares ZeRO-1/2 with ZeRO-0 after three steps. Fails on master (3.5e-2), passes here.
  • The step-0 check in test_a_discarded_step_leaves_every_momentum_where_it_was was vacuous. The buffer does not exist before the first step, so it compared None with 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/muon and tests/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.

_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

Copy link
Copy Markdown
Contributor Author

@pengdurice Thanks for setting auto-merge! It's still waiting on an approving review, in case that got missed.

@pengdurice
pengdurice added this pull request to the merge queue Sep 24, 2026
Merged via the queue into deepspeedai:master with commit 1355e9d Sep 24, 2026
13 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants