DeepSpeed
4a5856c2 - Run Muon once per optimizer step under ZeRO-3, not once per micro-batch (#8600)

Commit
17 days ago
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>
Author
Parents
Loading