Skip to content

Run Muon once per optimizer step under ZeRO-3, not once per micro-batch - #8600

Merged
delock merged 4 commits into
deepspeedai:masterfrom
alanhuangyoo:fix/zero3-muon-once-per-step
Sep 22, 2026
Merged

delock merged 4 commits into
deepspeedai:masterfrom
alanhuangyoo:fix/zero3-muon-once-per-step

Conversation

@alanhuangyoo

Copy link
Copy Markdown
Contributor

Fixes #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 #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 (Default gradient_clipping divides every Muon update by its own norm, shrinking the step by a model-sized factor #8439 / Fix Muon optimizer conflict with gradient clipping in ZeRO 1/2 #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 [muon] Keep the momentum out of steps the loss scaler discards #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; #8464 is working on it. jinyouzhi added a pointer to this shape in #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.

ZeRO-3 applied Muon inside the gradient reduce, which runs every micro-batch.
With gradient_accumulation_steps n, the momentum advanced n times per step and
Newton-Schulz orthogonalized partial gradients, so training with accumulation
diverged from the same batch taken in one micro-batch. ZeRO-1/2 were correct.

Without optimizer offload, the reduce path now leaves Muon out and the
partitions accumulate the raw gradient. At the step, after the overflow check
and before the norm, each Muon sub-group's accumulated gradients are gathered in
bounded chunks, orthogonalized once through the same round-robin update, and
written back to the partitions. The per-subgroup update is factored out of
_apply_distributed_muon_update so both paths share it. The offload path is
unchanged.

Fixes deepspeedai#8443

Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
@@ -0,0 +1,100 @@
# Copyright (c) Microsoft Corporation.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

should remove Microsoft copyright head.

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.

Removed, thanks.

m,
beta=self.muon_beta,
ns_method=getattr(self, 'muon_ns_method', 'gram'),
num_heads=getattr(param, 'muon_num_heads', None))

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

param is no longer defined, will always get a None here.

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.

Good catch, that was mine: pulling the grad from the list dropped the param = params[base_i + rank] line. Restored in 31f5e29. Added a test that takes one step with and without per-head under ZeRO-3 and checks it moves q and k but not the MLP. It fails without the fix.

Taking the gradient from the list dropped the line that bound `param` to
the parameter being updated, so `muon_num_heads` was read from whichever
parameter the momentum loop above left behind. Per-head Muon under ZeRO-3
then split every matrix in the bucket by one parameter's head count.

The new test takes one step with and without per-head from the same start:
per-head has to change q and k and leave the untagged MLP alone. With the
binding missing, the MLP picks up a head count and the test fails.

Also drops the Microsoft copyright line from the new test file.

Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
@alanhuangyoo
alanhuangyoo force-pushed the fix/zero3-muon-once-per-step branch from 31f5e29 to f494efb Compare September 21, 2026 17:22
Comment thread deepspeed/runtime/zero/stage3.py Outdated
numel += group[end].ds_numel
end += 1
params, chunk = group[start:end], partitions[start:end]
full_grads = self._partitioned_buffers_all_gather(params, chunk, chunk[0].dtype)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

This all-gather should use self.communication_data_type instead of chunk[0].dtype — otherwise a user-configured fp16 comm dtype is silently overridden and the gather's communication volume and buffer footprint double.

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.

Fixed in 952cc44, both gathers now use self.communication_data_type. Muon suite still passes (261 on 2 GPUs).

Comment thread deepspeed/runtime/zero/stage3.py Outdated
full_grads = self._partitioned_buffers_all_gather(params, chunk, chunk[0].dtype)
group_items = [(param, self.grad_position[self.get_param_id(param)][1], full_grad)
for param, full_grad in zip(params, full_grads)]
self._muon_update_sub_group(i, group_items, chunk[0].dtype)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Same for the momentum gather inside _muon_update_sub_group: the momentum buffer itself is stored in communication_data_type (_create_momentum_buffer), so gathering it in fp32 gains no precision — it only doubles the traffic. Please pass self.communication_data_type so both paths match the storage dtype.

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.

Same commit.

@delock

delock commented Sep 22, 2026

Copy link
Copy Markdown
Collaborator

Hi @alanhuangyoo , the updated PR looks overall good for me. The code review revealed a fix needed for communication data type, can you take a look at the comments? Thanks!

The step-time gather passed the partitions' own dtype, so a configured
fp16/bf16 communication dtype was ignored and the gradient and momentum
gathers moved twice the bytes. The momentum buffer is already stored in
the communication dtype, so gathering it wider gained nothing.

Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
@delock
delock enabled auto-merge September 22, 2026 11:07
The reduce path now calls _apply_distributed_muon_update unconditionally,
as it does on master, and the function returns early without optimizer
offload. Reading offload_optimizer at the call site broke
test_zero3_autoep_contiguous_grads_average_by_global_dp, which builds the
optimizer without __init__ and stubs that method out.

Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>
auto-merge was automatically disabled September 22, 2026 13:34

Head branch was pushed to by a user without write access

@alanhuangyoo

Copy link
Copy Markdown
Contributor Author

@delock The cpu-torch-latest failure was mine: test_zero3_autoep_contiguous_grads_average_by_global_dp builds the optimizer without __init__ and stubs out _apply_distributed_muon_update, and I had added an offload_optimizer check at the call site. Moved it inside the function in 3473e2c, so the reduce path is back to calling it unconditionally like on master. That test and the Muon offload/ZeRO-3 tests pass locally. The push turned auto-merge off, sorry, could you re-enable it?

@delock
delock enabled auto-merge September 22, 2026 14:53
@delock
delock added this pull request to the merge queue Sep 22, 2026
Merged via the queue into deepspeedai:master with commit 4a5856c Sep 22, 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.

ZeRO-3 applies Muon's Newton-Schulz once per micro-batch, so gradient accumulation changes the optimizer

2 participants