Skip to content

[muon] Keep the momentum out of steps the loss scaler discards - #8435

Merged
delock merged 4 commits into
deepspeedai:masterfrom
alanhuangyoo:fix/muon-momentum-survives-overflow
Sep 13, 2026
Merged

delock merged 4 commits into
deepspeedai:masterfrom
alanhuangyoo:fix/muon-momentum-survives-overflow

Conversation

@alanhuangyoo

Copy link
Copy Markdown
Contributor

Fixes #8432.

Problem

Muon under ZeRO 1/2 with fp16 does not train. Running the configuration tests/unit/ops/muon/test_muon.py itself uses, for longer than its five steps:

Exception: Current loss scale already at minimum - cannot decrease scale anymore. Exiting run.

No parameter is ever updated. The first loss-scale overflow is folded into Muon's momentum before the overflow check decides to discard the step, and the buffer stays non-finite for the rest of the run.

The failure sustains itself because both halves of muon_update touch the gradient:

momentum.lerp_(grad, 1 - beta)                     # inf/nan enters the momentum
update = grad.lerp_(momentum, beta) if nesterov    # ...and is written back into grad

so the next step's gradient is already non-finite whatever the loss scale has been reduced to. Backing off cannot help.

Fix

The momentum stays out of a step whose gradient is not finite, and the non-finite gradient is still returned so the overflow is seen and the step skipped.

Both halves are needed. Protecting the momentum alone is not enough: on the non-nesterov path update is the momentum, so a protected momentum would hand an overflowed step a finite update and the step would be applied rather than skipped. The non-finiteness has to keep propagating. Evaluated on device, so this costs no synchronization.

Verification

30 steps, 2 × H20, ZeRO 1/2, SimpleModel(hidden_dim=128, nlayers=5), lr=0.05:

stage initial scale master this branch
1 65536 (default) 0/10 moved, momentum non-finite, dies 10/10 moved, momentum finite
2 65536 (default) 0/10 moved, momentum non-finite, dies 10/10 moved, momentum finite
1 1 (no overflow) 10/10 moved, finite 10/10 moved, finite
2 1 (no overflow) 10/10 moved, finite 10/10 moved, finite

The scale-1 rows are the control: with no overflow there was never a problem, so the failure is entirely the overflow interaction rather than anything about Muon.

Tests

tests/unit/ops/muon/test_muon_overflow.py. Three of the four fail on the parent commit:

master   3 failed, 1 passed
           FAILED test_an_overflowed_gradient_does_not_enter_the_momentum
                  - AssertionError: an overflowed step must not move the momentum
           FAILED test_training_recovers_from_the_initial_overflow[1]
                  - Exception: Current loss scale already at minimum
           FAILED test_training_recovers_from_the_initial_overflow[2]
                  - Exception: Current loss scale already at minimum
this PR  4 passed

The unit-level case pins both halves directly: the momentum must not move, and the returned update must stay non-finite. test_a_finite_gradient_still_moves_the_momentum is the guard against the guard — it would catch a fix that simply disabled the optimizer.

Existing suite on this branch, non-offload configurations:

2 failed, 80 passed

Both failures are op_builder.builder.CUDAMismatchException in test_muon_reduce_scatter_with_optimizer_offload_raises, from this box's system CUDA not matching the one torch was built against, so CPUAdam will not build. They are unrelated to this change and reproduce on master.

Note on the suite

Worth recording, since it is why this survived: TestMuonConfigs captures initial_params before deepspeed.initialize, which casts the model to fp16. The assertion is therefore fp32 against fp16 and torch.equal is False whatever happened in between:

initial dtype (pre-init)            torch.float32
after-training dtype                torch.float16
repo assertion (pre-init vs after)  10/10 "changed"
same-dtype comparison               0/10 actually changed
cast alone, no training at all      10/10 "changed" by the same assertion

Casting a fresh model to fp16 and training it zero steps satisfies it. That is out of scope here — the new file captures after initialize and says why in a comment — but the assertion is worth tightening separately.

Muon under ZeRO 1/2 with fp16 does not train. The first loss-scale overflow is
folded into the momentum buffer before the overflow check decides to discard the
step, the buffer stays non-finite for the rest of the run, and every later step
overflows too until the scaler gives up:

    Exception: Current loss scale already at minimum - cannot decrease scale
    anymore. Exiting run.

No parameter is ever updated. This is the configuration test_muon.py itself uses,
only run for longer than five steps.

The failure sustains itself because both halves of muon_update touch the
gradient:

    momentum.lerp_(grad, 1 - beta)                    # inf/nan enters the momentum
    update = grad.lerp_(momentum, beta) if nesterov   # ...and is written back to grad

so the next step's gradient is already non-finite whatever the loss scale has
been reduced to.

The momentum now stays out of a step whose gradient is not finite, and the
non-finite gradient is still returned so the overflow is seen and the step is
skipped. Both are needed: on the non-nesterov path `update` is the momentum, so
protecting the momentum alone would hand an overflowed step a finite update and
the step would be applied instead of skipped. Evaluated on device, so this costs
no synchronization.

30 steps, 2 x H20, ZeRO 1/2, SimpleModel(hidden_dim=128, nlayers=5), lr 0.05:

    stage  scale     master                              this commit
    1      65536     0/10 moved, non-finite, dies        10/10 moved, finite
    2      65536     0/10 moved, non-finite, dies        10/10 moved, finite
    1      1         10/10 moved, finite                 10/10 moved, finite
    2      1         10/10 moved, finite                 10/10 moved, finite

The scale-1 rows are the control: with no overflow there was never a problem, so
the failure is entirely the overflow interaction.

Reported as deepspeedai#8432, which also records why the suite is green today: the run is
too short to reach the exception, and the parameter-change assertion compares
parameters captured before deepspeed.initialize -- fp32 -- against fp16 ones
after training, so torch.equal is False whatever happened in between. Casting a
model to fp16 and training it zero steps satisfies that assertion.

Tests: 3 of the 4 new cases fail on the parent commit, including the unit-level
one that pins the momentum directly.

Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
# scaler backs off until it raises "Current loss scale already at minimum". Keep the
# momentum out of it, and let the non-finite gradient through so the overflow is still
# seen and the step still skipped. Evaluated on device so this costs no synchronization.
grad_is_finite = torch.isfinite(grad).all()

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

grad_is_finite is per tensor, but the decision it protects the momentum from is global. has_overflow (stage_1_and_2.py:2482) sums _has_inf_or_nan over every partitioned gradient in every group and all-reduces MAX across the DP and model-parallel groups, and step() discards the whole step on that one flag at :2307. The ordering is not in doubt: the flag is computed from averaged_gradients, which is what get_flat_partition returns, and that is where muon_update is applied per tensor at :2167.

So on a step discarded because some other matrix overflowed, this matrix's gradient was finite, its momentum has already moved, and the parameter update it moved for is thrown away. The momentum then carries a step that never happened. That is narrower than the invariant your test module states ("A step the loss scaler discards must leave Muon's momentum as it was"), and it is the case the tests do not reach: both tensor-level tests use a single tensor, and the training test only asserts the momentum is finite, which the mixed case satisfies.

Is the per-tensor scope deliberate? One absorbed gradient is cheap next to the poisoned-forever behaviour you are fixing, so this is not a blocker either way. I read this rather than ran it, and I have no multi-GPU box here, so I have not seen the mixed case happen.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, deliberate — but you are right that the module claimed more than it delivers, and that is fixed in 4ce8e62. Apologies for not answering here; I pushed it and never replied in the thread.

Two changes, both from your reading:

The docstring no longer states the invariant you quoted. It is now "A tensor whose own gradient overflowed must not absorb it into its momentum", and it says the scope out loud — that the guard is per tensor, that has_overflow reduces over every partitioned gradient and step() discards on that one flag, and that a finite tensor therefore still advances its momentum on a step discarded for another.

And TestMuonMixedOverflow::test_a_finite_tensor_still_absorbs_a_step_discarded_for_another pins that case rather than leaving it to reading. You said you have no multi-GPU box, so here is the measurement — one finite matrix (calm) and one whose gradient overflows (boom), momentum read where muon_update writes it (inside get_flat_partition, during backward) rather than around engine.step():

this branch:  step 1  overflow=True  calm 55.159618 -> 104.801094   boom 42.780441 -> 42.780441  discarded
master:       step 1  overflow=True  calm 55.159618 -> 104.801094   boom 42.780441 -> inf
              step 2  overflow=True  calm 104.801094 -> 149.487137  boom inf -> nan
              step 3  overflow=True  calm 149.487137 -> 162.125839  boom nan -> nan

calm moves on a discarded step on both — that is the residue you describe, and it is the same on master, so this PR does not make it worse. What changes is the boom column: inf -> nan -> nan forever on master versus held at its pre-overflow value here.

On whether to close the residue too: it needs the momentum write deferred until the step is known to survive, which costs a second buffer the size of the momentum for every Muon parameter. That did not seem worth trading for one absorbed gradient per overflow event, so the docstring records the choice instead of hiding it. Happy to be argued out of that if you or a maintainer would rather pay the memory.

@alanhuangyoo

Copy link
Copy Markdown
Contributor Author

You are right, and it was not deliberate — I put the guard where muon_update already was and did not think about the flag being global.

Tracing it the same way you did, the ordering is forced, not incidental:

  • self.averaged_gradients[i] = self.get_flat_partition(...) (:956, :989) — muon_update runs inside get_flat_partition (:2167), so the momentum moves at gradient-reduction time.
  • has_overflow_partitioned_grads_serial (:2474) reads self.averaged_gradients[i], i.e. after that.
  • step() discards on the result at :2307.

So the global flag is computed from post-muon_update gradients by construction. There is no ordering of the current code that lets a per-tensor call see it.

What a global guard would actually cost, since that is the part your comment leaves open:

  1. Check overflow before applying Muon. Means an extra all-reduce inside the reduction path, on every step, to buy a rollback that only matters on discarded steps. Wrong trade at fp16 loss-scale frequencies.
  2. Invert the update. momentum.lerp_(grad, 1 - beta) inverts exactly in real arithmetic ((m_new - (1-beta)*g) / beta), and in the mixed case this tensor's grad is finite so it is well-defined — but it is not bit-exact in fp32, so "leave the momentum as it was" would become "leave it approximately as it was".
  3. Defer the write. Compute the update from a copy and commit the momentum only once the step is known to survive. Exact, no extra collective, costs one buffer the size of the momentum — which is already a full copy of the Muon parameters, so roughly +1× on those groups.

(3) is the only one that delivers what my test module claims. Whether that memory is worth it for a case that costs one absorbed gradient is a call I would rather you and @delock make than assume.

Two things I am doing either way:

  • The invariant in the test module is wrong as written. "A step the loss scaler discards must leave Muon's momentum as it was" is stronger than a per-tensor guard delivers. I will narrow it to what is actually true — a tensor whose own gradient overflowed does not absorb it — and add the mixed case explicitly, so the gap is recorded rather than implied.
  • Measure it. You said you read this rather than ran it and have no multi-GPU box; I do. Two 2-D parameters, only one fed an overflowing input, fp16 + ZeRO-1, and check whether the finite one's momentum moves on the discarded step. My box is unreachable right now, so this is a promise rather than a result — I will post the numbers, including if they contradict the reading.

Thanks for reading it this closely. This is the second time on this stack I have claimed something wider than I measured, and both times it was someone else who noticed.

@alanhuangyoo

Copy link
Copy Markdown
Contributor Author

Numbers, as promised. Your case reproduces exactly.

Two 2-D parameters in one group, calm and boom. Only boom is fed an input that overflows in fp16, at step 1. Muon, fp16, ZeRO-1, single GPU. The momentum is read from the group's flat partition buffer before and after each backward, because that is where muon_update writes — not inside step().

On this branch:

step  overflow   calm momentum (before bwd -> after bwd)   boom momentum (before -> after)   params
   0     False          None -> 55.159618                       None -> 42.780441            changed
   1      True     55.159618 -> 104.801094                 42.780441 -> 42.780441            NO (discarded)
   2     False    104.801094 -> 149.487137                 42.780441 -> 81.290314            changed
   3     False    149.487137 -> 189.690079                 81.290314 -> 115.945297           changed

Step 1 is discarded and the parameters do not move, boom is protected — and calm moves anyway, 55.159618 -> 104.801094, for an update that is thrown away. That is your case, and it is not hypothetical.

On master, same script:

step  overflow   calm momentum                            boom momentum                     params
   0     False          None -> 55.159618                       None -> 42.780441            changed
   1      True     55.159618 -> 104.801094                 42.780441 -> inf                  NO (discarded)
   2      True    104.801094 -> 149.487137                        inf -> nan                 NO (discarded)
   3      True    149.487137 -> 162.125839                        nan -> nan                 NO (discarded)

So the same run shows both things at once: the poisoning this PR is for (inf -> nan, every later step discarded, loss scale halving 16 -> 8 -> 4, no recovery), and the gap you found (calm advancing on a discarded step, identically on both sides — the fix does nothing for it).

What I am changing here: the invariant in the test module, which as written promises more than the code delivers, and a test that pins the mixed case at the behaviour above so it is recorded rather than discovered again.

What I am not changing without a word from you and @delock: the scope. Of the three ways to make it global, only deferring the momentum write is exact, and it costs one buffer the size of the momentum. One absorbed gradient per discarded step against +1x optimizer memory on the Muon groups is a trade I would rather not make unilaterally inside a PR that is meant to stop a run from dying.

Thanks — you found this by reading, without a box to run it on, and you were right on every detail including which tests could not reach it.

@ebarkhordar

Copy link
Copy Markdown
Contributor

Thanks for running it. Your step 1 row is the case exactly: calm moving 55.159618 to 104.801094 on a step whose parameter update is discarded, and identical on both sides, so the guard does not touch it.

On scope, since you asked. Deferring the momentum write is the only one of the three that delivers what the test module claims. I would avoid inverting the lerp for the reason you gave, that it turns an exact invariant into an approximate one, and a rollback bought with an extra collective on every step is paying on the common path for the rare one. Whether the buffer is worth one absorbed gradient per discarded step is a memory call I have no numbers for, so that part is yours and @delock's.

Narrowing the invariant and pinning the mixed case is worth doing either way.

The module claimed 'A step the loss scaler discards must leave Muon's momentum
as it was'. The guard is per tensor and the scaler's decision is global --
has_overflow reduces _has_inf_or_nan over every partitioned gradient and step()
discards on that one flag -- so a tensor whose own gradient was finite still
advances its momentum on a step discarded for another tensor.

Measured on 1xH20, two 2-D parameters in one group, only one fed an overflowing
input, fp16 + ZeRO-1, momentum read on either side of backward because that is
where muon_update writes:

  step  overflow  calm momentum              boom momentum            params
     0     False        None -> 55.159618          None -> 42.780441   changed
     1      True   55.159618 -> 104.801094    42.780441 -> 42.780441   discarded

On master the same run gives boom 42.780441 -> inf, then nan, with every later
step discarded and the loss scale halving to the minimum -- the failure this PR
is for. calm's 55.159618 -> 104.801094 is identical on both sides.

Reported by @ebarkhordar, who found it by reading and named which tests could
not reach it: both tensor-level cases use a single tensor and the training test
only asserts the momentum is finite.

Docstring now says what the guard covers, and
test_a_finite_tensor_still_absorbs_a_step_discarded_for_another records the gap
so a later change to global scope shows up as a failing test rather than a
silent improvement.

Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
@alanhuangyoo

Copy link
Copy Markdown
Contributor Author

Done, in 4ce8e62 — the uncontested half.

The module docstring now says what the guard covers (a tensor whose own gradient overflowed does not absorb it) instead of what it does not (a step the scaler discards leaving the momentum as it was), and spells out why the two differ.

test_a_finite_tensor_still_absorbs_a_step_discarded_for_another pins the case: two 2-D parameters in one group, only one fed an overflowing input, fp16 + ZeRO-1, momentum read on either side of backward. It asserts that the overflowing tensor's momentum holds, that no parameter moves, and that the finite tensor's momentum does advance — with a message saying that if this starts failing, the guard has become global and the docstring should follow.

So the gap is a failing test away from being noticed rather than a paragraph someone has to remember.

tests/unit/ops/muon/test_muon_overflow.py: 3 passed, 2 skipped (fp16 world_size=2 cases need the second GPU).

Leaving the scope decision where you put it — @delock, the question is whether one absorbed gradient per discarded step is worth a second buffer the size of the momentum on the Muon groups. I have no preference strong enough to spend your memory budget on it.

@delock

delock commented Sep 12, 2026

Copy link
Copy Markdown
Collaborator

Hi @alanhuangyoo @ebarkhordar , thanks for the thorough discussion. Is that true if we take option 3, the additional buffer should be temporary? How likely it would contribute to peak GPU memory?

My preference lean to correctness. make weight not updated but momentum updated sounds an inconsistency that should not have. A second buffer is the least complex solution to me. I would suggest to leave a memory profile measurement result (peak memory before/after), in this PR comment section as a record for future reference if people wants to optimize the memory part.

alanhuangyoo and others added 2 commits September 12, 2026 23:51
@delock chose correctness, so this takes option 3 and closes the half a per-tensor
guard cannot reach: a tensor whose own gradient was finite no longer advances its
momentum on a step the loss scaler discards.

`muon_update` still runs while the partition is filled, but it now writes into a
staging buffer that starts as a copy of the committed momentum. `step()` commits it
once the overflow check has passed. A discarded step simply never reaches the commit,
so the momentum is left bit-identical rather than approximately so -- which is why this
is option 3 and not the lerp inversion.

The staging buffers live on the optimizer rather than in `optimizer.state`, so they do
not travel into checkpoints, and they are reused across steps rather than reallocated:
staging re-copies from the committed momentum every step, so a stale one cannot be
committed.

`test_a_finite_tensor_still_absorbs_a_step_discarded_for_another` pinned the gap and
said "if this starts failing, the guard has become global". It has, so it is now
`test_a_discarded_step_leaves_every_momentum_where_it_was` and asserts the whole
invariant, including that a surviving step still moves both momenta. Its reads moved
from around `backward` to around the whole step, since the commit is in `step()` now.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
@alanhuangyoo

Copy link
Copy Markdown
Contributor Author

@delock Option 3 is in (b1223e4), and the memory profile is below as a record.

Answering the two questions directly

"Should the additional buffer be temporary?" It can be, but it does not buy what you would expect. Measured both ways:

peak allocated resident after the step staging held
master 2715.6 MiB 1089.6 MiB —
this PR (buffer reused across steps) 2971.6 MiB 1345.6 MiB 256.0 MiB
same, freed after each commit 2971.6 MiB 1089.6 MiB 0

SimpleModel(hidden_dim=4096, nlayers=8), 134.2M Muon parameters, fp16, ZeRO-1, 1×H20, four steps, torch.cuda.max_memory_allocated().

"How likely would it contribute to peak GPU memory?" It contributes exactly its own size, +256.0 MiB = 1× the momentum buffer, and freeing it after each commit does not change that. The peak occurs while the buffer is alive — muon_update runs inside the reduction epilogue, and the Newton-Schulz working set is allocated on top of it — so returning the memory afterwards lowers the resident figure and not the peak.

So the trade is: reused costs 1× momentum resident and nothing on the allocator path; temporary returns that between steps at the cost of an allocate/free per step, with the same peak. I went with reused, since the peak is what usually decides whether a run fits. Flipping it is a one-line change if you would rather have the resident memory back.

Both variants are correct — I ran the behaviour check against each.

The behaviour

Same harness as before, fp16 + ZeRO-1, one injected overflow at step 2:

                 master                         this PR
step  overflow   momentum before -> after       momentum before -> after
   0     False    0.000000 -> 40.940689          0.000000 -> 20.691717
   1     False   40.940689 -> 65.035875         20.691717 -> 35.152515
   2      True   65.035875 ->       nan         35.152515 -> 35.152515   <- discarded
   3      True         nan ->       nan         35.152515 -> 40.018685
   4      True         nan ->       nan         40.018685 -> 43.899609
   5      True         nan ->       nan         43.899609 -> 46.261569

master is #8432 in six lines: one overflow, then every subsequent step overflows and the scaler backs off to its minimum. On this branch the discarded step leaves the momentum bit-identical (35.152515 both sides) and training continues.

On the test that had to flip

test_a_finite_tensor_still_absorbs_a_step_discarded_for_another pinned the gap and carried this note:

if this starts failing, the guard has become global and the module docstring should say so

It has, so it is now test_a_discarded_step_leaves_every_momentum_where_it_was and asserts the whole invariant — plus that a surviving step still moves both momenta, so a future change cannot satisfy it by disabling the optimizer. Its reads moved from around backward to around the whole step, because the commit is in step() now.

tests/unit/ops/muon/test_muon_overflow.py: 5 passed.

@ebarkhordar — this is the half you said was a memory call rather than yours to make; the numbers above are what it costs.

@delock
delock enabled auto-merge September 13, 2026 04:10
@delock
delock added this pull request to the merge queue Sep 13, 2026
Merged via the queue into deepspeedai:master with commit b5e000c Sep 13, 2026
13 checks passed
yh0903 pushed a commit to yh0903/DeepSpeed that referenced this pull request Sep 22, 2026
…ch (deepspeedai#8600)

Fixes deepspeedai#8443.

## The problem

ZeRO-3 applies Muon inside the gradient reduce
(`_apply_distributed_muon_update`, called from
`__avg_scatter_contiguous_grads`), and that runs every micro-batch. With
`gradient_accumulation_steps: n`, the momentum advances `n` times per
optimizer step, and Newton-Schulz orthogonalizes each micro-batch's
partial gradient instead of the accumulated one. ZeRO-1/2 apply Muon at
the accumulation boundary and are correct.

On 2 GPUs, fp32, with the same 8 samples per step either way (one
micro-batch of 8 at `gas=1`, four of 2 at `gas=4`), three steps,
relative difference in the weights:

| | master | this PR |
| --- | ---: | ---: |
| ZeRO-2, `gas=1` vs `gas=4` | 3.6e-4 | 3.6e-4 |
| ZeRO-3, `gas=1` vs `gas=4` | **1.3e-1** | 3.6e-4 |
| Newton-Schulz calls, ZeRO-3, 2 matrices, 2 steps, `gas=4` | 16 | 4 |

ZeRO-3 now lands on exactly ZeRO-2's figure.

## The change

This is option 1 from the discussion in deepspeedai#8443. It is scoped to ZeRO-3
without optimizer offload.

- The reduce path no longer runs Muon when optimizer offload is off. The
partitions accumulate the raw averaged gradient, as they do for every
other optimizer.
- `step()` calls `_apply_muon_to_accumulated_grads()` after the overflow
check and before the gradient norm. For each Muon sub-group, it:
- all-gathers each parameter's accumulated gradient partitions, in
chunks bounded by `reduce_bucket_size` as the reduce buckets were;
  - runs the existing round-robin Muon update once;
  - writes each rank's slice back into its partition.
- The per-sub-group body of `_apply_distributed_muon_update` is moved
into `_muon_update_sub_group`. It takes the full-shape gradients
explicitly, because at step time the parameters are partitioned and
`param.grad` can't hold them. The reduce path calls it with `param.grad`
as before.
- The gradient norm is still taken over the Muon update, as before.
Clipping semantics are unchanged (deepspeedai#8439 / deepspeedai#7776 are separate).
- Because the update now runs after the overflow check, a step the loss
scaler discards no longer touches the momentum. That is the ZeRO-3
counterpart of deepspeedai#8435.
- Collectives: each Muon parameter's gradient and momentum are gathered
once per step instead of once per micro-batch. At `gas=n` that is `n`
times fewer.

The optimizer-offload path is unchanged; deepspeedai#8464 is working on it.
jinyouzhi added a pointer to this shape in deepspeedai#8464, and the overlap is
limited to `_apply_distributed_muon_update`.

## Testing

On 2×H20:

- New file `tests/unit/v1/ops/muon/test_muon_zero3_grad_accum.py`. Both
tests fail on master and pass here.
- `test_newton_schulz_runs_once_per_matrix_per_step`: Newton-Schulz
calls summed over ranks come to 2 × steps, not 2 × steps × gas.
- `test_gradient_accumulation_matches_one_large_micro_batch[2, 3]`:
`gas=1` and `gas=4` agree to within half-precision Newton-Schulz noise
at both stages.
- `tests/unit/v1/ops/muon/` plus
`tests/unit/runtime/zero/test_per_head_muon.py`: 260 passed.

---------

Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
pull Bot pushed a commit to AmirulAndalib/DeepSpeed that referenced this pull request Sep 24, 2026
deepspeedai#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 deepspeedai#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 (deepspeedai#8439), and that norm is taken over
different tensors at each stage.

Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
Co-authored-by: pengdurice <pengduhit@gmail.com>
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.

[BUG] Muon + fp16 does not train under ZeRO 1/2: the first loss-scale overflow permanently poisons the momentum buffer

3 participants