Repository navigation
[Perf] MAGI-2: drop the Ulysses send/receive buffer copies - #8595
Draft
yeahdongcn wants to merge 4 commits into
Draft
yeahdongcn wants to merge 4 commits into
yeahdongcn wants to merge 4 commits into
Conversation
…uffer The packed Q/K/V exchange copied every input into a destination-rank-major temporary and then concatenated the three temporaries into the send buffer, moving the whole buffer twice per layer. Each input is now copied once, directly into its head slice of a destination-rank-major send buffer, so the concatenation pass goes away. The send buffer has the same shape, dtype, layout and bytes as the concatenated one, and the receive buffer and the per-input views that FlashAttention reads are unchanged; only the order of the copies differs. Inputs on different devices are rejected up front instead of by the concatenation. Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
After the output all-to-all, each rank copied its [world, S_rank, H, D] receive buffer to [S_rank, world*H, D] before the gated output projection. On MUSA the attention kernel now hands back the [S_rank, world, H, D] head-shard view of the receive buffer, and Magi2Attention.output flattens it after the token gather, so inside the compiled region the gather reads the receive buffer directly and the eager copy goes away. Only the layout of the compiled region's attention input changes: every value reaching the gate multiply, and so the output projection input, is the same element of the same receive buffer. scatter_seqlen_gather_heads keeps its contract and the other platforms keep the flattened output. Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
On MUSA, Magi2Attention.project returns Q/K/V as [S_rank, world, H_i, D] head-shard views of one [world, S_rank, sum(H_i), D] buffer, so Q/K/V are laid out in all-to-all send order. scatter_heads_gather_seqlen recognises views that tile one destination-rank-major buffer and sends that buffer as is; any other input, including a view that autograd tracks, still takes the single copy into a fresh send buffer, and a single-rank exchange flattens head shards back to [S_rank, H, D]. The Q/K/V values are computed by the same operations as before and only stored at different addresses: the head-shard layout is a pure permutation of the old send buffer, whose bytes, split sizes and receive-side views are unchanged. The layout is used only when every Ulysses rank gets the same number of query and KV heads. Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
This was referenced Oct 7, 2026
…tion Under dynamic-shape compilation every input dimension of a compiled region is symbolic, and Inductor passes symbolic sizes to Triton as 64-bit arguments. The attention-gate kernel that gathers the [T, world, H/world, D] head-shard view of the receive buffer therefore splits each flat index with 64-bit divisions and remainders by the row size H*D, the head size D, the head count H and the per-rank head count H/world, and the gated product is copied once more before the output projection. When the head shards come from the module's own Ulysses group, Magi2Attention.output specializes their shard count and checks the per-rank head count and the head size with torch._check. A dynamic-shape compile then keeps only the token count symbolic: the gate kernel indexes the receive buffer with constants and the output projection reads its result directly. Head shards of any other shard count take the unpinned path. Static compiles and eager runs are unaffected, and the gate arithmetic is unchanged, so the outputs are bitwise identical in both compile modes. Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Purpose
Each MAGI-2 attention layer runs two Ulysses all-to-alls, and both carried extra full-buffer copies. Before the Q/K/V all-to-all, every input was copied into a destination-rank-major temporary and the three temporaries were concatenated into the send buffer, so the whole buffer moved twice. After the output all-to-all, each rank copied its
[world, S_rank, H, D]receive buffer to[S_rank, world*H, D]before the gated output projection. This PR removes those copies in three commits and keeps the result fast under dynamic-shape compilation with a fourth:[Perf]Copy Ulysses Q/K/V straight into the all-to-all send buffer (all platforms).scatter_heads_gather_seqlenallocates one[world, S_rank, sum(H_i), D]send buffer and copies each input once into its head slice, so the concatenation pass goes away. The send buffer has the same shape, dtype, layout and bytes as before, and the receive buffer and the per-input views FlashAttention reads are unchanged. Inputs on different devices are now rejected up front.[Perf]Read the Ulysses attention output in place on MUSA. The newscatter_seqlen_gather_head_shardsreturns the[S_rank, world, H, D]view of the receive buffer. On MUSA the attention kernel hands back that view, andMagi2Attention.outputflattens it after the token gather, so the compiled output region reads the receive buffer directly.scatter_seqlen_gather_headskeeps its contract.[Perf]Store Ulysses Q/K/V straight into the send buffer on MUSA. On MUSA,Magi2Attention.projectreturns Q/K/V as[S_rank, world, H_i, D]head-shard views of one destination-rank-major buffer (pack_ulysses_head_shards), so Q/K/V are laid out in send order andscatter_heads_gather_seqlencan send that buffer without a copy. Any other input, including a view autograd tracks, takes the single copy from commit 1.[Perf]Pin the Ulysses head-shard layout in the output projection. Withdiffusion_compile_dynamic=True, the default, every input dimension of a compiled region is symbolic and Inductor passes symbolic sizes to Triton as 64-bit arguments, so the attention-gate kernel that reads the[T, world, H/world, D]view split each flat index with four 64-bit divisions and remainders, and the gated product was copied once more before the output projection. When the head shards come from the module's own Ulysses group,Magi2Attention.outputnow specializes the shard count and checks the per-rank head count and the head size withtorch._check, so only the token count stays symbolic and the kernel indexes with constants; head shards of any other shard count take the unpinned path. Static compiles and eager runs are unchanged.Commit 1 is the only change to the data path on other platforms. Commits 2 and 3 change the layouts inside compiled regions, so they are enabled on MUSA only, decided once in
Magi2Attention.__init__; the exchange helpers they extend keep their 3-D contracts. Commit 3 also requires the query and KV head counts to divide by the Ulysses degree; with MAGI-2's 24 heads that holds at SP8 (3 heads per rank) and SP4 (6 heads per rank). No change alters a value: Q/K/V and the attention output are computed by the same operations and are only stored at different addresses, and the generated video and audio stay byte-identical (Test Result).Based on main and independent of the other open MAGI-2 PRs. The overlaps are textual: with #8511 the
.parallelimport list inmodeling_magi2.py, and with the per-request rope/modality and sink-LSE caches PR three adjacent lines inulysses_packed_attention_with_sinkandMagi2PackedAttentionKernel; whichever lands second needs a trivial rebase. Running MAGI-2 on MUSA needs #8498.Test Plan
python -m pytest -q -o addopts='' \ tests/diffusion/models/magi2/test_ulysses_exchange.py \ tests/diffusion/models/magi2/test_native_distributed_parity.py \ tests/diffusion/models/magi2/test_native_compile_distributed.pytest_ulysses_exchange.pyruns both exchanges for every rank of 2-, 4- and 8-rank groups in one process, with uneven splits, a rank without tokens, unequal head counts and mixed dtypes. It compares the send buffers, split sizes and receive views bit for bit with the concatenating implementation, checks that packed head shards are sent in place and that inputs which do not tile one buffer are copied, checks the MUSA gating, and checks thatMagi2Attention.projectandMagi2Attention.outputgive bitwise-identical results with and without the head-shard layouts. On a MUSA device it compilesMagi2Attention.outputandMagi2TransformerLayer._attention_inputstatically withemulate_precision_castsat production model widths with 3702 tokens (8 ranks) and 3651 (4 ranks), plus a 14-token case, compares every output bit for bit with the unpacked layout, and asserts that the compiled region hands back views of one send buffer. No test counts kernels, so that the compiled producers store straight into that buffer, rather than Inductor adding a copy into it, is shown only indirectly, by the one-run end-to-end gain of the three commits together. These tests compile with Inductor'sdeterministicmode andtest_configs.force_filter_reduction_configs, so two compiles of the same code pick the same reduction configs. For commit 4, a CPU test tracesMagi2Attention.outputwith dynamic shapes and checks that only the token dimension of matching head shards stays symbolic, and that a mismatched shard count still compiles and gives the same result; on a MUSA device a test compiles the output projection dynamically at world sizes 8 and 4 with one and three modalities and asserts one gather kernel without symbolic divisions, the same number of Triton kernels as the static compile, and output bitwise equal to the static compile.The 4-rank gloo parity test adds an SP4 run with the MUSA layouts forced on CPU and requires it to equal plain SP4 bit for bit. The distributed compile test adds the same layouts as an SP2 variant traced with
fullgraph=True,backend="eager"and dynamic shapes undererror_on_recompile, and requires its compiled run to equal its eager run and its eager run to equal plain SP2 bit for bit; on MUSA it needs two torchada fixes and passes with them (Test Result).End to end: MAGI-2 Preview T2VA, prompt "A golden retriever running through a sunlit meadow, cinematic camera movement", 272p (448x256), 125 frames, 100 steps, seed 42,
diffusion_compile_dynamic=False, on 8x MTT S5000. Configurations (--expert-parallel-sizeis from #8511):--ulysses-degree 8 --cfg-parallel-size 1 --enable-expert-parallel --expert-parallel-size 4--ulysses-degree 4 --cfg-parallel-size 2--ulysses-degree 4 --cfg-parallel-size 2 --enable-expert-parallel --expert-parallel-size 4The per-step figures come from an out-of-tree timing harness around
AsyncOmniwith the per-step boundary from #7274; the metric is the per-stepsampler.diffusetime, maximum over the 8 ranks, mean over steps 2-99.Byte-identity check: Inductor benchmarks reduction configs on every fresh compile, so the outputs of one code base differ across fresh caches, and
TORCHINDUCTOR_DETERMINISTIC=1does not reach the compiled regions: the MAGI-2 pipeline callstorch.use_deterministic_algorithmsat construction, which resets Inductor's deterministic flag. The fixed-config comparisons therefore ran withTORCHINDUCTOR_DETERMINISTIC=1 TORCHINDUCTOR_FORCE_FILTER_REDUCTION_CONFIGS=1and a freshTORCHINDUCTOR_CACHE_DIRper run, and compare the sha256 of the generated video and audio. GEMM shape padding is still chosen by benchmarking under this setting, so differing hashes alone would not show a change in numerics; identical hashes do show identical output.vLLM Version: 0.28.0 (MUSA image
vllm:v0.28.0-ph1-5.2.0-torch2.11.0.post2-20261001, torch 2.11.0.post2, torch_musa 2.11.0.post2+musa5.2.0), with a local test-only shim for the vLLM 0.30 names that c1e84ce imports; current main aligns with vLLM 0.31, which was not run. The single-change A/B ran with torchada c045bd3 (MooreThreads/torchada#120); the device suites and the fixed-config byte-identity runs also had torchada'scpp_extensionpath-signature fix MooreThreads/torchada#121 (host-side build paths only), without which cold-cache CPU Inductor compiles fail on MUSA.vLLM-Omni Commit: 5a9da74 on top of c1e84ce, the base of the other MAGI-2 PRs; no MAGI-2 file has changed on main since. The tests and runs used the same three changes on c1e84ce, where their code is identical apart from neighbouring lines of other changes (the
.parallelimport list, and the sink-LSE arguments next to the new attention parameter). Trees named in Test Result:Test Result
test_native_compile*.pywith the FXNameError: name 'device' is not definedand fail identically on plain main c1e84ce in this image. The per-test list was not kept. By the failure count, the 4 are the four tests that compile regions withdynamic=True, one of which istest_native_compile_distributed.py, which this PR extends; its new MUSA-layout variant therefore did not run in that suite; see the commit 4 Test Plan entry below. The gloo parity test, which runs the same layouts eagerly, was not among the failures.TORCHINDUCTOR_DETERMINISTIC=1with a fresh cache but did not filter reduction configs; the output hashes matched the base runs, and byte identity rests on the fixed-config check below.diffusion_compile_dynamic=True, the default), SP8xCFG1+EP4: control 761.00 ms; integrated 785.71 ms; integrated without [Perf] MAGI-2: drop the Ulysses send/receive buffer copies #8595 741.57 ms; without the second commit of [Perf] MAGI-2: use FP32 mHC stream contractions on MUSA #8510 807.50 ms; without [Perf] MAGI-2: run the MUSA mHC Sinkhorn iterations in one Triton kernel #8594 759.46 ms. Without [Perf] MAGI-2: run the MUSA mHC Sinkhorn iterations in one Triton kernel #8594 the output is identical to the dynamic control, so every other change is byte-identical under dynamic compile too; with [Perf] MAGI-2: run the MUSA mHC Sinkhorn iterations in one Triton kernel #8594 the output equals the static-compile output (see [Perf] MAGI-2: run the MUSA mHC Sinkhorn iterations in one Triton kernel #8594). Only runs with identical outputs are compared for time: the eager MoE GEMM time follows the routing, which changes with the output bits.test_musa_pipeline.py, 2setup_compiletests intest_native_compile.py), 4 are thetest_native_compile*.pyFXNameError(a torchada issue in itstorch.devicereplacement, fixed in fix(device): register the device factory as the FX device builtin MooreThreads/torchada#124), 1 istest_expert_parallel.py::test_ep[tencent/HunyuanImage-3.0](downloads a config), and 8 aretest_omni_config.py::test_diffusion_stage_payload_keys_roundtrip, which fail the same way on main.diffusion_compile_dynamic=Truethe first three commits cost about 44 ms per step at SP8xCFG1+EP4 (785.71 with them vs 741.57 without, identical outputs), because the attention-gate kernel that reads the head-shard view (triton_poi_fused_index_select_mul_sigmoid_view) went from 56.8 to 106.3 ms/step; that kernel is already about 50x slower under dynamic than under static shapes on main (fixed separately in [Perf] MAGI-2: pin the attention head layout in the output projection #8627). Commit 4 removes it; see the next bullet._dense_outputregion: dynamic compile without commit 4 2628 µs per call, with seven 64-bit divisions (ks0-ks4typedi64in the Triton signature) and a copy kernel; with commit 4 39.0 µs, no division, no copy kernel, equal to the static compile (39.1 µs), output bitwise identical.test_ulysses_exchange.py165 passed (143 CPU-side, 22 MUSA).test_native_compile_distributed.pyneeds two torchada fixes on MUSA: the FXdevicebuiltin (fix(device): register the device factory as the FX device builtin MooreThreads/torchada#124, merged), and FSDP2 on a CPU mesh (fix(fsdp): run torch_musa's FSDP2 replacements only on MUSA meshes MooreThreads/torchada#126; torch_musa's FSDP2 replacements always create MUSA streams, so the test's gloo CPU-mesh HSDP fails in FSDP's root pre-forward). With both,test_compiled_regions_match_eager_under_dlo_and_hsdppasses on MUSA at 5a9da74, including the new MUSA-layout variant, with MUSA devices visible and hidden. In the same runtest_native_compile.pypasses apart from the twosetup_compiletests, which need network access.AI assistance: Claude Code drafted the change, the tests and this description and ran the MUSA validation listed above.