Skip to content

[Perf] MAGI-2: pin the attention head layout in the output projection - #8627

Open
yeahdongcn wants to merge 1 commit into
vllm-project:mainfrom
yeahdongcn:xd/magi2-dynamic-attention-gate-pins
Open

yeahdongcn wants to merge 1 commit into
vllm-project:mainfrom
yeahdongcn:xd/magi2-dynamic-attention-gate-pins

Conversation

@yeahdongcn

@yeahdongcn yeahdongcn commented Oct 8, 2026 •

Copy link
Copy Markdown
Contributor

Purpose

With diffusion_compile_dynamic=True, the default, every input dimension of a compiled MAGI-2 region is symbolic, and Inductor passes symbolic sizes to Triton as 64-bit arguments. The attention-gate kernel in Magi2Attention.output, which gathers the [T, H, D] attention output in modality order and multiplies it by sigmoid(gates), therefore splits every flat index with 64-bit divisions and remainders by the row size H*D and the head size D, and the gated product is copied once more before the output projection. On MTT S5000 at 3702 tokens per rank this kernel takes 1.43 ms per call under dynamic compile, plus 0.21 ms for the copy, against 55 µs under static compile; MAGI-2 calls it 40 times per step.

Magi2Attention.output now checks the head layout of its input with torch._check: the dimensions between the token and head-size dimensions hold num_heads_q heads (attention.shape[1:-1].numel()), and the last dimension is head_dim. A dynamic-shape compile specializes them, so only the token count stays symbolic; the gate kernel indexes with constants, as under static compile, and the output projection reads its result directly. The gate arithmetic is unchanged, and static compiles and eager runs are unaffected (the static generated code is identical), so the outputs are bitwise identical in both compile modes.

The head count is checked over all those dimensions so that the check also holds when the heads are split over two dimensions, as in the [T, world, H/world, D] head shards that #8595 hands to output() on MUSA. The two PRs therefore work in either landing order; they edit the same first lines of output(), so whichever lands second needs a trivial rebase.

Based on main and independent of the other open MAGI-2 PRs.

Test Plan

python -m pytest -q -o addopts='' tests/diffusion/models/magi2/test_native_compile.py

The new test traces Magi2Attention.output with dynamic=True for two token counts, checks that one graph serves both and that only the token dimension of its attention input is symbolic, and checks that running the traced graph gives the eager output bit for bit. Without the checks it fails on the head-layout assertion. It does not run Inductor; the Inductor-level evidence is the end-to-end comparison below.

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, SP8xCFG1 with EP4 (--ulysses-degree 8 --cfg-parallel-size 1 --enable-expert-parallel --expert-parallel-size 4, from #8511), on 8x MTT S5000, with diffusion_compile_dynamic=True and with False. The per-step figures come from an out-of-tree timing harness around AsyncOmni with the per-step boundary from #7274; the metric is the per-step sampler.diffuse time, maximum over the 8 ranks, mean over steps 2-99. Byte identity compares the sha256 of the generated video and audio. Inductor picks reduction configs and GEMM padding by benchmarking on each fresh compile, which can change output bits, and TORCHINDUCTOR_DETERMINISTIC=1 does not reach the compiled regions because the pipeline resets it at construction; the runs therefore set TORCHINDUCTOR_FORCE_FILTER_REDUCTION_CONFIGS=1 with a fresh TORCHINDUCTOR_CACHE_DIR per run. 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. torchada with MooreThreads/torchada#120 and #121 (both since merged).

vLLM-Omni Commit: 8c79740 on top of c1e84ce; it merges cleanly into main 4c5541c, which has not changed modeling_magi2.py or test_native_compile.py since. The end-to-end runs add this commit to c1e84ce with #8496, #8497, #8498, #8506, the first commit of #8510 (all five PRs since merged) and the first three commits of #8511 (the control).

Test Result

  • CPU: test_native_compile.py passes apart from the two setup_compile tests, which import vLLM and could not run in the local CPU environment. With [Perf] MAGI-2: drop the Ulysses send/receive buffer copies #8595 applied on top, its head-shard tests (test_ulysses_exchange.py) also pass, and compiled dense and head-shard inputs still give bitwise-identical outputs at 2, 4 and 8 shards under static and dynamic compile.
  • MUSA, Test Plan: the new test passes. Of the other six tests one passes and five fail, as they do without this change: three with NameError: name 'device' is not defined, a torchada issue in its torch.device replacement, fixed in fix(device): register the device factory as the FX device builtin MooreThreads/torchada#124 (merged), and the two setup_compile tests, which fetch a config from the Hugging Face Hub while the test machine had no network access.
  • Attention-gate kernel as Inductor generates it on MUSA, 3702 tokens, L2 flushed before every call, median over 5 rounds: dynamic compile 1429.3 µs per call plus a 212.2 µs copy kernel; with this change 55.2 µs and no copy; static compile 55.1 µs. The dynamic kernel's Triton signature types its four size arguments as i64; with the checks only the token count remains. The outputs of the three forms are bitwise identical.
  • Dynamic compile, SP8xCFG1 with EP4: 756.83 -> 696.21 ms per step (-60.6 ms, -8.0%), video and audio sha256 identical. The saving matches the kernel figures above (40 calls per step).
  • Static compile, SP8xCFG1 with EP4: 713.58 -> 714.80 ms per step, sha256 identical.
  • Not run: CUDA, where the same specialization applies; it changes only which sizes Inductor treats as constants.

AI assistance: Claude Code drafted the change, the test and this description and ran the MUSA validation listed above. I reviewed every change in this PR and re-ran the validation listed above.

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, H, D]
attention output in modality order therefore splits each flat index
with 64-bit divisions and remainders by the row size H*D and the head
size D, and the gated product is copied once more before the output
projection.

Magi2Attention.output checks the head count and the head size of its
attention input with torch._check. A dynamic-shape compile specializes
them, so only the token count stays symbolic: the gate kernel indexes
with constants and the output projection reads its result directly.
The head count is checked over all dimensions between the token and
head-size dimensions, so the check also holds when the heads are split
over two dimensions. 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>
@chatgpt-codex-connector

Copy link
Copy Markdown

Codex usage limits have been reached for code reviews. Please check with the admins of this repo to increase the limits by adding credits.
Credits must be used to enable repository wide code reviews.

@vllm-omni-review-bot

Copy link
Copy Markdown

This PR appears to belong to: docs/design/module/diffusion/index.md, docs/design/module/diffusion/offloader.md, docs/design/module/diffusion/diffusion_model_integration.md.

Module owners: @wtomin @RuixiangMa @david6666666

Routing: @wtomin via module of the changed files, module named in the PR description, semantic router, CODEOWNERS; @RuixiangMa via module of the changed files, module named in the PR description, semantic router; @david6666666 via module of the changed files, module named in the PR description

@yeahdongcn, please review your own changes and leave a short self-review comment describing what you checked. PRs without author self-review may not be assigned a reviewer.

Please take a look when you have a chance. If you would like an automated review, mention @vllm-omni-review-bot in a comment.

@vllm-omni-review-bot

Copy link
Copy Markdown
Omni ReviewBot routing record

Assigned Strict on cursor (cursor-grok-4.6-high) under experiment fleet-strict-cursor-grok46-zcode-glm53flash-5050-c5-z10-20261002.

@vllm-omni-review-bot vllm-omni-review-bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Omni ReviewBot review

0 actionable finding(s).

CI at 8c79740c67b0 (2026-10-10T02:54:40.733326+00:00): required check(s) blocking: buildkite/vllm-omni (missing).

Full review analysis

Scan:

Category Result
Tests / verification no finding reported
Security no finding reported
Docs / comments no finding reported
Behavior / compatibility no finding reported
Correctness no finding reported

Validated:

  • [resolved] #8595 4D [T,world,H/world,D]: product check holds so landing-order assert will not fire; residual: it does not specialize world and H/world, and 4D attention × 3D gates [T,H,1] does not broadcast — #8595 must still edit output()
  • [resolved] Eager/static compile: _check is a tautology on the THD layout production always feeds; existing test_compile_regions_fullgraph_matches_eager_bitwise would fail if 3D tiny-model output() were rejected. Residual: 4D [T,world,H/world,D] from unlanded #8595 is not a valid full output() input here because permute * sigmoid(gates) still assumes 3D gates.
  • [resolved] #8595 4D [T,world,H/world,D] satisfies shape[1:-1].numel(); residual: 3D gates multiply still cannot broadcast that 4D layout — author said the second-landing PR rebases output().
  • [resolved] #8595 4D [T,world,H/world,D]: product check still holds; residual is permute×3D gates mis-broadcast — current gather is still 3D at parallel.py:235 reshape(local_tokens, group.world_size * tensor.shape[1], tensor.shape[2]); rebase assigned to whichever PR lands second.
  • [claim-verified] 3702 tokens/rank matches 272p packing: width,height=448,256; VAE stride (8,16,16); internal 250 frames → latent 32×16×28=14336 video; audio=250; CFG packed ×2; SP8 even split ⇒ text length 222 ((2*(14336+250+222))/8=3702).
  • [validated] modeling_magi2.py:194-195 torch._check(attention.shape[1:-1].numel() == self.num_heads_q) and torch._check(attention.shape[-1] == self.head_dim) pin local TP heads; only caller _mlp_input at modeling_magi2.py:574; attend() is [T,H,D] (attention.py:113); Ulysses gather reshapes back to 3D at parallel.py:235.

No actionable findings.

Verdict: APPROVE

No actionable findings.


🤖 This review was generated by InferMatrix Copilot, an open-source repo-maintenance agent for PR review, CI debugging and issue triage. Try it on your own repo, and ⭐ star it if it helped!

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants