Run Muon once per optimizer step under ZeRO-3, not once per micro-batch (#8600)
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 (#8439 / #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 #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.
---------
Signed-off-by: alanhuangyoo <alanhuangyoo@gmail.com>