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>