DeepSpeed
5880d38c - Support Muon in BF16_Optimizer with FP32 gradient accumulation (#8526)

Commit
14 days ago
Support Muon in BF16_Optimizer with FP32 gradient accumulation (#8526) BF16 weights with FP32 gradient accumulation and ZeRO stage 1 select `BF16_Optimizer`. Muon currently refuses that configuration because this wrapper hands the optimizer flat partitions without first applying Newton–Schulz. This change computes Muon updates on the original matrices after gradient reduction and global clipping, then copies only each rank's owned intersections into the optimizer update. Momentum remains FP32 and partitioned like the master weights. The implementation gathers one group's momentum as temporary workspace, stages local momentum until the optimizer step, and excludes alignment padding. Auxiliary Adam groups retain their existing update path. The initial support boundary is dense data parallelism with two-dimensional Muon matrices; model-parallel MPU, expert groups and graph harvesting remain rejected. This adds communication and temporary memory proportional to the largest Muon group; it is a correctness change, with no performance claim. Validation on two NVIDIA L20 GPUs (PyTorch 2.13.0+cu130): | Test | Configuration / coverage | Result | |---|---|---| | BF16 wrapper regressions | Three pytest cases, each exercised at world sizes 1 and 2; standard/Gram NS, accumulation boundaries, partition-sized FP32 momentum, and the neighboring BF16-gradient configuration | **3 passed** | | Eager engine E2E | DP=2, accumulation over 3 microbatches, clipping=0.05, mixed Muon/Adam groups, and 17×13 / 9×17 matrices crossing partition boundaries; standard and Gram NS | All 3 optimizer steps match the independent full-matrix reference **exactly** | | Checkpoint continuation | Save, continue for 2 steps, reload, and replay; standard and Gram NS | Losses, weights, and optimizer state reproduce **exactly** | | Final source validation | Repeat the eager E2E using the final source files in an isolated directory | **Passed**, including exact reference comparison and checkpoint replay | | Compiled NS | Same-input, per-step FP32-master comparison with upstream compilation enabled | **Passed** at `atol=1e-4, rtol=1e-3`; checkpoint replay remains exact | | Repository checks | All configured pre-commit hooks on the 4 changed files | **Passed** | The compiled comparison is a same-input per-step check. Independently evolving BF16 trajectories can diverge after a rounding difference; these results do not establish bitwise compiled equivalence. Repository regression command: ```bash TORCHDYNAMO_DISABLE=1 PYTHONPATH=.:tests python -m pytest -q \ tests/unit/runtime/zero/test_muon_without_zero_optimizer.py \ -k TestMuonBF16Optimizer ``` This addresses the missing orthogonalization in `BF16_Optimizer`. It is separate from #8483's ZeRO-1/2 gradient/momentum dtype reconciliation and #7748's checkpoint dtype conversion. Fixes #8461. --------- Signed-off-by: 0z5a <0z5a@users.noreply.github.com> Signed-off-by: 0z5a <dezhen.lu@student.uni-tuebingen.de> Co-authored-by: 0z5a <0z5a@users.noreply.github.com> Co-authored-by: Ma, Guokai <guokai.ma@gmail.com>
Author
Parents
Loading