Skip to content

[Perf] MAGI-2: use FP32 mHC stream contractions on MUSA - #8510

Merged
hsliuustc0106 merged 2 commits into
vllm-project:mainfrom
yeahdongcn:xd/magi2-musa-mhc-contractions
Oct 9, 2026
Merged

hsliuustc0106 merged 2 commits into
vllm-project:mainfrom
yeahdongcn:xd/magi2-musa-mhc-contractions

Conversation

@yeahdongcn

@yeahdongcn yeahdongcn commented Oct 5, 2026 •

Copy link
Copy Markdown
Contributor

Purpose

On MUSA, MAGI-2's four-stream mHC contractions lower to small-K batched GEMMs: the pre-application and the stream mix contract over four streams, and the logits projection is one skinny [T, 4C] x [4C, 24] GEMM. On MTT S5000 these are the largest remaining mHC cost. On MUSA only:

  • MHCHandler.apply_pre and the compiled stream mix use an FP32 elementwise product and sum, which Inductor fuses with the neighbouring pointwise ops.
  • MHCHandler.compute_logits splits the projection by stream into a batched FP32 matmul plus a sum.

The MUSA branch is decided once in MHCHandler.__init__. Other platforms keep the existing einsum and matmul forms, and vllm_omni/diffusion/layers/mhc.py is unchanged: eager MUSA runs still use the Triton mix kernel from #7261, and the Sinkhorn post-residual path (#7545) is not touched. apply_pre takes the FP32 form only for exact-shaped [T, N] coefficients with the stream dtype (fp32, bf16 or fp16) outside autocast; broadcast, single-stream, float64, mixed-dtype and autocast inputs keep the einsum path, including its errors.

Precision, measured on MTT S5000 at 4096 tokens x 4 streams x 1024 channels:

  • Eager: products of BF16/FP16 inputs are exact in FP32, so only the order of the four-term sum changes. Over 99.98% of the outputs are bit-identical to the einsum form, with the same maximum error against FP64.
  • torch.compile: Inductor keeps the fused coefficients in FP32 rather than rounding them to the stream dtype first, so about 40% of the outputs differ in the last bit. Their error against the FP64 product with FP32 coefficients equals the einsum form's error against its own reference.
  • The per-stream projection's maximum error against FP64 drops from 1.25e-5 to 4.7e-6.
  • End to end the generated video differs from the einsum form (PSNR 24.9 dB, audio 37.4 dB). Switching unchanged code from SP4xCFG2 to SP8xCFG1 gives a similar difference (24.0 dB, 38.0 dB), and side-by-side frames show the same scene.

The second commit, [Perf] MAGI-2: store the MUSA mHC norm output stream-major, removes a copy in front of that batched matmul. torch_musa's bmm copies an operand whose batches are interleaved before its GEMM, and the [streams, tokens, hidden] view of the flat [T, 4C] norm output is interleaved (strides (C, 4C, 1)), so every compute_logits call copied the whole FP32 operand with a CopyLastContiguousKernel of about 0.3 ms at 272p. An explicit .contiguous() does not help under Inductor: the clone is fused into the norm reduction and keeps the reduction's token-major layout. Under compilation, MHCHandler._stream_major_normed now runs the norm on the [tokens, streams, hidden] view and copies the result into a stream-major buffer behind an as_strided, which pins the buffer's strides, so the fused norm kernel stores the dense [streams, tokens, hidden] layout itself and bmm reads it in place. The reduction loop of that kernel is unchanged and, in the generated code we inspected, a producing hyper-connection mix stays fused into it, so the outputs are bit-identical (see Test Result for how this was measured). On ranks whose tokens span several modalities, the multimodal layers still need one copy, now made by Inductor instead of bmm. Eager MUSA runs make one explicit copy. For this, MultiModalityRMSNorm accepts features split over its last two dimensions; it reduces over both and views the weight as [streams, hidden], and a single-dimension input traces to the same ops as before. Other platforms are unchanged.

Part of the MAGI-2 MUSA work split from #7156; mHC mixing is in Section 3 of #7085. Running MAGI-2 on MUSA needs #8498. The second commit's A/B below runs SP8xCFG1 with head-sharded EP4, which needs #8511; the change itself does not depend on it.

Test Plan

python -m pytest -q -o addopts='' \
  tests/diffusion/models/magi2/test_mhc_stream_contractions.py \
  tests/diffusion/layers/test_mhc.py \
  tests/diffusion/models/magi2/test_native_preview.py

The new tests check that, off MUSA, all three contractions are bit-equal to the einsum/matmul forms; that the MUSA forms stay within an FP32 error bound of an FP64 reference and do not round the products before the sum; that the einsum contract (broadcasting, errors, float64) is kept outside the fast path; and, on a MUSA device, the eager and torch.compile results.

The second commit adds CPU tests that the split-feature norm matches the flat norm for single- and multi-modality layouts and rejects inputs that do not split its features, and that the eager and compiling compute_logits hand bmm a dense stream-major operand and return exactly the old logits. On a MUSA device, a test compiles the old strided-operand form and the new form at hidden 3072 (1031 and 3702 tokens, one to three modality groups, with and without a fused hyper-connection mix), checks in the generated code that bmm receives the dense stream-major buffer, and compares the outputs bit for bit. The module compiles with Inductor's deterministic mode and filtered reduction configs, since otherwise each compile benchmarks its own reduction config and the two forms could pick different ones.

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, SP4xCFG2, diffusion_compile_dynamic=False, on 8x MTT S5000. The second commit was measured at SP8xCFG1 with EP4 (--ulysses-degree 8 --cfg-parallel-size 1 --enable-expert-parallel --expert-parallel-size 4), and the integration tree carrying it was also checked for byte identity at SP4xCFG2 (--ulysses-degree 4 --cfg-parallel-size 2) and SP4xCFG2 with EP4. 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 of the second commit: Inductor benchmarks reduction configs on every fresh compile, so the output of one code base differs across fresh caches, and TORCHINDUCTOR_DETERMINISTIC=1 does not reach the compiled regions: the MAGI-2 pipeline calls torch.use_deterministic_algorithms at construction, which resets Inductor's deterministic flag. The comparison therefore ran with TORCHINDUCTOR_DETERMINISTIC=1 TORCHINDUCTOR_FORCE_FILTER_REDUCTION_CONFIGS=1 and a fresh TORCHINDUCTOR_CACHE_DIR per run, and compares 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), with a local test-only shim for the vLLM 0.30 names that main imports. The A/B and the profile of the second commit ran with torchada MooreThreads/torchada#120; the fixed-config byte-identity runs and the test suites also had torchada's cpp_extension path-signature fix MooreThreads/torchada#121 (host-side build paths only), without which cold-cache CPU Inductor compiles fail on MUSA.

vLLM-Omni Commit: ae9bce4 (second commit) on top of a07b481 (first commit) on top of c1e84ce. The first commit was tested at a07b481. The second commit was tested in an integration tree with the other pending MAGI-2 changes, where its hunks are identical.

Test Result

AI assistance: Claude Code drafted the change, the tests and this description and ran the MUSA validation listed above.

On MUSA the four-stream mHC contractions lower to small-K batched GEMMs,
which are slow. On MUSA only, apply_pre (for same-dtype [T, N]
coefficients outside autocast) and the compiled stream mix use an FP32
elementwise product and sum, which Inductor fuses, and compute_logits
splits the projection by stream into a batched FP32 matmul plus a sum.
Eager MUSA runs keep the Triton mix kernel. Other platforms are
unchanged.

Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
@vllm-omni-review-bot

Copy link
Copy Markdown

This PR appears to belong to: docs/design/module/diffusion/offloader.md, docs/design/module/diffusion/diffusion_model_integration.md, docs/design/module/diffusion/index.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.

@hsliuustc0106 hsliuustc0106 added high priority high priority issue, needs to be done asap ready/cuda-test Reviewed and ready for CUDA testing enhancement New feature or request labels Oct 6, 2026
@Dong1017

Dong1017 commented Oct 7, 2026

Copy link
Copy Markdown
Contributor

I found no additional correctness issue in this revision. The MAGI-2 call path supplies FP32 mHC normalization and projection tensors; the new pre-application guard preserves the einsum path outside its exact supported subset, and eager MUSA mixing retains the existing Triton implementation. The existing CPU contract run recorded 18 passed and 3 MUSA cases deselected; it was not rerun in this double-check. MUSA device behavior, S5000 performance, and model quality remain author-reported. The PR describes the timing harness as out-of-tree but does not link its source or invocation, so I could not reproduce that comparison here.

torch_musa's bmm copies an operand whose batches are interleaved before
its GEMM, so handing it the [streams, tokens, hidden] view of the flat
mHC norm output costs a full copy per call. An explicit .contiguous()
does not help under Inductor, which keeps the clone in the token-major
order of the norm reduction it fuses into. Under compilation the norm
now reduces over the [tokens, streams, hidden] view and writes into a
stream-major buffer whose layout an as_strided consumer pins, so the
norm kernel stores that layout itself. The generated reduction is
unchanged and a producing hyper-connection mix stays fused into it, so
with fixed reduction configs the outputs are bit-identical. Eager makes
one explicit copy. Other platforms are unchanged.

MultiModalityRMSNorm accepts features split over the last two
dimensions.

Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>

@hsliuustc0106 hsliuustc0106 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Re-reviewed at ae9bce4. The MUSA-only contractions and stream-major norm layout check out, including native fallbacks and compile behavior. No new correctness findings. Static checks passed; device tests and performance measurements were not independently rerun.

@hsliuustc0106 hsliuustc0106 added the ready label to trigger buildkite CI label Oct 9, 2026 — with ChatGPT Codex Connector
@hsliuustc0106
hsliuustc0106 merged commit 4c5541c into vllm-project:main Oct 9, 2026
7 of 9 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request high priority high priority issue, needs to be done asap ready/cuda-test Reviewed and ready for CUDA testing ready label to trigger buildkite CI

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants