[Performance] Partition Muon 2D params per dtype to restore comm fusion (#79740)
* [Performance] Partition Muon 2D params per dtype to restore comm fusion
_partition_2d_parameters decided each 2D param's owner rank by greedy
bin-packing over all dtypes at once, sizing active_ranks from the total
volume of the whole color group. A FusedCommBuffer never mixes dtypes,
because AssignGroupBySize keys its groups on dtype, so a dtype holding
only a small share of the params (e.g. a handful of fp32 params among
bf16 ones) got spread over as many ranks as the whole group needed. Each
of those ranks then built its own tiny buffer of that dtype and the
reduce degenerated into many small-payload calls.
Run the same greedy packing once per dtype so each dtype occupies only
as many ranks as its own volume requires. dtypes are walked in a sorted
order so every rank derives an identical mapping, otherwise the reduce
dst and the broadcast src would disagree across ranks.
The bin-packing itself, the `* 4` byte estimate and the
`int(size / buffer) + 1` rounding are all kept as-is, so a single-dtype
parameter list produces exactly the previous mapping. With mixed dtypes
the owner mapping does change; the math is unaffected (each 2D param is
still owned and Newton-Schulz'd by exactly one rank) but gradient reduce
roots move, so results are not bit-identical to before.
* Add regression tests for per-dtype Muon 2D owner partitioning
Two levels of coverage for _partition_2d_parameters, plus one comment fix.
test/collective/fleet/test_muon_sharding_mixed_dtype_partition.py (2 cards):
builds bf16 and fp32 2D params in the same color group -- bf16 via
amp.decorate(level='O2'), fp32 kept out through excluded_layers -- with
comm_buffer_size_MB=1 and shapes sized so the two partitioning strategies
disagree. Asserts that all ranks derive the same owner mapping, that bf16
spans both ranks while fp32 concentrates on one, that the fp32 params fuse
into a single comm buffer, and that parameters remain bit-identical across
ranks after real optimizer steps. A regime guard fails the test if the
partition degenerates to the trivial single-rank case, which is what the
existing Muon cases silently do (small model + default 256 MB bucket).
test/legacy_test/test_muon_2d_partition.py (single process, no accelerator):
drives the method against stub params over a world_size x bucket-size x
param-shape matrix. Pins that a single-dtype list reproduces the previous
whole-group mapping exactly, that params are neither dropped nor duplicated,
that each dtype spreads only as wide as its own volume requires, that owner
assignment does not depend on dtype discovery order, and that the caller's
list is not reordered.
test/collective/fleet/CMakeLists.txt is hand-edited rather than regenerated:
gen_ut_cmakelists.py currently aborts on this directory because 9 of the 84
rows in testslist.csv name test files that no longer exist. The added block
mirrors the neighbouring Muon entry and takes the next free dist UT port.
Also corrects the comment on the dtype sort. It claimed the sort keeps the
reduce dst and the broadcast src aligned across ranks, which is not what it
does: every dtype is packed from rank 0 with a fresh size vector, so dtype
order cannot move a param to another rank. What it pins is the order of each
rank's param list, which is what _local_2d is built from.
* Fix mixed-dtype Muon test on devices without bfloat16
The 2-card case failed on CI: every 2D param came out float32, so the regime
guard fired with "the two dtypes were not both present among 2D params".
Cause is in amp_decorate (python/paddle/amp/auto_cast.py): on a GPU place it
returns the models untouched when the requested dtype is unsupported -- bf16
needs Compute Capability >= 8 -- and reports nothing. The distributed fleet
cases run on V100 (CC 7.0), where that branch is taken and the requested bf16
cast never happens, leaving a single-dtype color group and nothing for the test
to distinguish.
Rather than predicting the device rule, the low-precision candidates are now
applied in order (bfloat16, then float16) and the parameter dtype read back, so
the test uses whichever the device actually accepted; fp16 only needs CC >= 7.
Every assertion keys on the dtype that was applied instead of hardcoding bf16.
The shape choices are unaffected: the partitioner estimates every dtype's
volume as numel * 4, so 4 x [512, 512] still spans two ranks and
2 x [256, 256] still fits on one whichever low-precision dtype is in play.
Verified on 2 GPUs both ways, with the bf16 gate left alone and with it forced
off to reproduce V100:
bf16 path: buffers per dtype={'bfloat16': 2, 'float32': 1}, fp32 owners=[0]
fp16 path: buffers per dtype={'float16': 2, 'float32': 1}, fp32 owners=[0]
and confirmed the case still discriminates on the fp16 path -- with the old
whole-group partitioning it fails there too, with fp32 spread over both ranks
as two 0.25 MB buffers.
* Keep the mixed-dtype Muon test in range under float16
Two holes found while verifying the float16 fallback added in the previous
commit.
The weights were U[0, 1) and unscaled, so four 512-wide layers multiplied the
activation scale by ~dim/2 each: 147 -> 3.6e4 -> 9.4e6 -> 2.4e9. That is fine in
bf16 and fp32 but overflows float16 at the third layer, and the run this test
now takes on V100 is float16.
Worse, it overflowed silently. An inf loss yields nan parameters, and
np.testing.assert_array_equal treats nan as equal to nan, so the closing
"every rank holds identical parameters" check passed while testing nothing --
which is exactly why the earlier local float16 run looked green.
Weights are now randn/sqrt(dim), which holds the activation max near 3 through
all four layers in both float16 and bfloat16, and two assertions close the
silent-overflow path: the loss must be finite at every step, and the gathered
parameters must be finite before they are compared across ranks.
Verified on 2 GPUs: bfloat16 and the forced-float16 (V100) path both pass; the
old whole-group partitioning still fails the case on the float16 path, so it
keeps its discriminating power; and restoring the unscaled init on the float16
path fails with "step 0 loss is inf, not finite", so the new guard is live
rather than decorative.
* Drop the duplicated _short_dtype in the mixed-dtype Muon test
The helper was defined twice with identical bodies; the second definition
shadowed the first, so the two could drift apart under a later edit without
anything failing. Keep the one that precedes its first use in _build_model.
No behaviour change; the 2-card case still passes on both the bfloat16 and the
forced-float16 (V100) path. A
AlAuAu committed
df1fbe1dd882ee4b82bb5d9741a1357e27dc51e6
Parent: 64ef1c7
Committed by GitHub <noreply@github.com>
on 9/5/2026, 2:18:40 AM