Repository navigation
[Perf] MAGI-2: use FP32 mHC stream contractions on MUSA - #8510
Conversation
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>
|
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. |
Omni ReviewBot routing recordAssigned Strict on cursor (cursor-grok-4.6-high) under experiment |
|
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
left a comment
There was a problem hiding this comment.
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.
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_preand the compiled stream mix use an FP32 elementwise product and sum, which Inductor fuses with the neighbouring pointwise ops.MHCHandler.compute_logitssplits 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, andvllm_omni/diffusion/layers/mhc.pyis unchanged: eager MUSA runs still use the Triton mix kernel from #7261, and the Sinkhorn post-residual path (#7545) is not touched.apply_pretakes 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:
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 second commit,
[Perf] MAGI-2: store the MUSA mHC norm output stream-major, removes a copy in front of that batched matmul. torch_musa'sbmmcopies 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 everycompute_logitscall copied the whole FP32 operand with aCopyLastContiguousKernelof 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_normednow runs the norm on the[tokens, streams, hidden]view and copies the result into a stream-major buffer behind anas_strided, which pins the buffer's strides, so the fused norm kernel stores the dense[streams, tokens, hidden]layout itself andbmmreads 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 ofbmm. Eager MUSA runs make one explicit copy. For this,MultiModalityRMSNormaccepts 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.pyThe 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.compileresults.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_logitshandbmma 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 thatbmmreceives 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 aroundAsyncOmniwith 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 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=1does not reach the compiled regions: the MAGI-2 pipeline callstorch.use_deterministic_algorithmsat construction, which resets Inductor's deterministic flag. The comparison therefore ran withTORCHINDUCTOR_DETERMINISTIC=1 TORCHINDUCTOR_FORCE_FILTER_REDUCTION_CONFIGS=1and a freshTORCHINDUCTOR_CACHE_DIRper 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'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: 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
test_native_compile*.pywithNameError: name 'device' is not defined, which fail identically on plain c1e84ce in this image.TORCHINDUCTOR_DETERMINISTIC=1and a fresh Inductor cache per run: 677.38 ms per step with it vs 703.14 / 700.34 without it (-24.4 ms, -3.5%). The output hashes matched the base runs; byte identity rests on the fixed-config check below, since that setting alone does not fix the reduction configs.CopyLastContiguousKernelcalls that copied the[4, T, 3072]FP32bmmoperand, 80 per step and about 24 ms per step, are gone.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: itstorch.devicereplacement; to be fixed there), 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.AI assistance: Claude Code drafted the change, the tests and this description and ran the MUSA validation listed above.