Repository navigation
[Perf] MAGI-2: bit-exact MoE route build, token-major routes and separate shared-expert activation - #8592
[Perf] MAGI-2: bit-exact MoE route build, token-major routes and separate shared-expert activation#8592yeahdongcn wants to merge 3 commits into
Conversation
The BF16 fused MoE path spends part of every MoE layer on route bookkeeping and on layout copies around its two grouped GEMMs. This removes that work without changing any value on MUSA. _align_bf16_routes takes the sorted expert ids from the sort it already runs and finds each expert's bounds with searchsorted on them, so the atomic scatter_add_ histogram is gone; one cumsum of the padded counts serves the destinations, the block experts and the padded count. _bf16_fused_moe_forward numbers routes token-major (token * heads + head), so the W13 GEMM reads x_heads as it is and the W2 result is summed straight into the [tokens, heads] layout, without the input permute copy and the copy back to token order. The router builds its FP32 bmm operand head-major in one cast+permute copy instead of casting and then letting MUSA bmm copy the strided operand. For the same route ids the alignment metadata is unchanged. With token-major ids every expert block holds the same routes, and the router logits, top-k indices and MoE output are bit-identical on MUSA. test_moe_routing_parity.py checks each step against reference copies of the histogram, head-major and strided-operand forms, on CPU and on device. Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
The MoE layer concatenated the shared and modality-specific fc1 outputs, applied SwiGLU7 once and split the result. Inside the compiled region the concatenation turns both fc1 GEMMs into writes to column slices of one buffer and leaves the shared fc2 input strided, and the MUSA GEMM copies such operands and outputs. SwiGLU7 pairs adjacent columns and each fc1 output has an even width, so applying it to each projection is the same math, and on MUSA the outputs are bit-identical; both GEMM outputs and both down-projection inputs are now dense. test_shared_experts.py compares the result with the concatenated activation bit for bit: on CPU in BF16 (the native expression that compiled regions trace), and on MUSA in eager mode and under torch.compile. Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
The compiled MUSA tests compare two separately compiled forms bit for bit. Their fixture set torch._inductor.config.deterministic once, but Dynamo restores torch.use_deterministic_algorithms after every traced frame, and that also resets the Inductor flag. Only the first compile of each test ran in deterministic mode; the other form, and each recompile for a new group layout, could pick GEMM shape padding by benchmarking and, on MUSA, rescale reduction blocks, so the two forms could be lowered differently. Pass deterministic=True in the torch.compile options instead. Inductor applies them to every compile and recompile, so both forms are lowered under the same configuration. The assert message names the group layout, and the Dynamo reset runs even when a comparison fails. Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>
0feba33 to
b5fe47f
Compare
|
Codex usage limits have been reached for code reviews. Please check with the admins of this repo to increase the limits by adding credits. |
Omni ReviewBot routing recordAssigned Strict on zcode (GLM-5.3-Flash) under experiment |
|
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, docs/design/module/benchmarking.md. Module owners: @Bounty-hunter @wtomin @RuixiangMa Routing: @Bounty-hunter via module of the changed files, module named in the PR description, semantic router, model owner, CODEOWNERS; @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 @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 attempt recordReview attempt ended as failed (step 'review' (agent.review_diff): unhandled error: RuntimeError: zcode exited 1 without a result event: statusCode: undefined } Error: Turn execution failed (traceId: b818f66a-7510-4fa9-987e-4e8b6a5b7fdc) — check |
Omni ReviewBot attempt recordReview attempt ended as failed (step 'review' (agent.review_diff): unhandled error: RuntimeError: zcode exited 1 without a result event: statusCode: undefined } Error: Turn execution failed (traceId: 61460401-e15b-4f23-a4e8-004f7671610e) — check |
Omni ReviewBot attempt recordReview attempt ended as failed (step 'review' (agent.review_diff): unhandled error: RuntimeError: zcode exited 1 without a result event: statusCode: undefined } Error: Turn execution failed (traceId: 2ef92438-d25a-48a6-9445-c70d8765df73) — check |
vllm-omni-review-bot
left a comment
There was a problem hiding this comment.
Omni ReviewBot review
0 actionable finding(s).
CI at
b5fe47fb8332(2026-10-10T03:11:40.361326+00:00): required check(s) blocking:buildkite/vllm-omni(missing).
Note: The assigned review arm
strict/zcode/GLM-5.3-Flashcould not complete this review, so it was produced by the fallback armdirect/cursor/auto. It is excluded from the routing experiment.
Full review analysis
PR description
MAGI-2's BF16 multi-head MoE builds expert-block metadata from the sort it already runs: searchsorted replaces the int32 scatter_add_ histogram, and one padded cumsum fills destinations, block experts, and the padded count. Routes are numbered token-major so the grouped GEMMs read [tokens, heads, hidden] and write that layout back, and the router casts its BMM operand to contiguous head-major FP32 in the same copy. Shared-expert SwiGLU7 now runs on each fc1 projection alone, so both fc1 outputs and both fc2 inputs stay dense; with even widths that is the same pairing as activating the concatenation.
Change flow
flowchart LR
acts["[EXISTING] Token-head activations"]:::existing
align["[CHANGED] searchsorted route bounds"]:::changed
gemm["[CHANGED] Token-major BF16 MoE GEMM"]:::changed
shared["[CHANGED] Separate shared-expert SwiGLU7"]:::shared
tests["[NEW] Bit-exact parity tests"]:::new
out["[EXISTING] MoE and shared-expert outputs"]:::existing
acts --> align --> gemm --> out
acts --> shared --> out
align --> tests
gemm --> tests
shared --> tests
classDef existing fill:#e5e7eb,stroke:#6b7280,color:#111827
classDef changed fill:#fef3c7,stroke:#d97706,color:#451a03,stroke-width:2px
classDef new fill:#dcfce7,stroke:#16a34a,color:#052e16,stroke-width:2px
classDef removed fill:#fee2e2,stroke:#dc2626,color:#450a0a,stroke-width:2px
class shared changed
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!
Purpose
Two bit-exact changes to MAGI-2's multi-head MoE layer that take route bookkeeping and layout copies out of every MoE layer. Two commits, plus a test fix:
[Perf] MAGI-2: bit-exact BF16 MoE route build and layout(mh_moe.py):_align_bf16_routestakes the sorted expert ids from the sort it already runs and finds each expert's bounds withsearchsorted, so the atomicscatter_add_histogram is gone. One cumsum of the padded counts serves the destinations, the block experts and the padded count._bf16_fused_moe_forwardnumbers routes token-major (token * heads + head). The W13 GEMM readsx_headsas it is and the W2 result is summed straight into the[tokens, heads]layout, so the input permute copy and the copy back to token order are gone._routebuilds the FP32 router bmm operand head-major in one cast+permute copy. On MUSA, bmm otherwise makes its own contiguous copy of the strided operand.[Perf] MAGI-2: activate the shared-expert projections separately(modeling_magi2.py): SwiGLU7 pairs adjacent columns and both fc1 outputs have even widths, so activating each projection on its own is the same math as activating their concatenation. Both fc1 GEMMs now write dense outputs inside the compiled region, and both fc2 GEMMs read dense inputs, instead of the column slices that the MUSA GEMM copies.[Test] MAGI-2: compile the shared-expert comparisons deterministically(test_shared_experts.py): the compiled MUSA tests passdeterministic=Truein thetorch.compileoptions, so both forms are lowered under the same configuration in every compile (see Test Result).No value changes on MUSA: for the same route ids the route metadata is identical, every expert block holds the same routes, and the router logits, top-k indices, MoE output and shared-expert output are bit-identical there. End to end, the generated video and audio are byte-identical (see Test Result). On CPU in FP32 the old concatenated form can round 1 ulp differently (strided sigmoid input, column-slice fc2 operand), so the CPU shared-expert check runs in BF16.
This PR applies to main, and its tests use nothing from another open PR.
_align_bf16_routesno longer needs the int32scatter_add_histogram from #8496 and drops it. The head-major reference intest_moe_routing_parity.pyrepeats main's BF16 GEMM launch config, including the 64-wide K tile on MUSA from #8506.Part of the MAGI-2 MUSA work split from #7156.
Test Plan
python -m pytest -q -o addopts='' \ tests/diffusion/models/magi2/test_moe_routing_parity.py \ tests/diffusion/models/magi2/test_shared_experts.py \ tests/diffusion/models/magi2/test_bf16_moe_wiring.py \ tests/benchmarks/test_magi2_bf16_moe_benchmark.pyThe new tests compare each step bit for bit with reference copies of the code it replaces (the routing references are copies of main's
_align_bf16_routesand_bf16_fused_moe_forwardbefore this PR, with main's launch config; the shared-expert reference, on CPU and MUSA, hands fc2 dense halves rather than main's column slice, and main's exact form is covered on MUSA by the end-to-end byte identity):torch.compile, plus SwiGLU7 alone at production widths.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. Three configs: SP8xCFG1 with head EP4 (--ulysses-degree 8 --cfg-parallel-size 1 --enable-expert-parallel --expert-parallel-size 4, from #8511), SP4xCFG2 (--ulysses-degree 4 --cfg-parallel-size 2), and SP4xCFG2 with head 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: Inductor benchmarks reduction configs on every fresh compile, so the output bits 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 comparison runs therefore setTORCHINDUCTOR_DETERMINISTIC=1 TORCHINDUCTOR_FORCE_FILTER_REDUCTION_CONFIGS=1with 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), with a local test-only shim for the vLLM 0.30 names that c1e84ce, the base of these runs, imports (main now targets vLLM 0.31). torch 2.11.0.post2, torch_musa 2.11.0.post2+musa5.2.0, torchada c045bd3 (MooreThreads/torchada#120, merged). The fixed-config byte-identity runs and the test suites also had torchada'scpp_extensionpath-signature fix MooreThreads/torchada#121 (merged; host-side build paths only), without which cold-cache CPU Inductor compiles fail on MUSA.vLLM-Omni Commit: b5fe47f on main 4c5541c, which contains #8496, #8497, #8498, #8506 and #8510. The device results below were taken on c1e84ce-based trees. The measurements used a stack of c1e84ce with #8496, #8497, #8498 and #8506, plus #8510 and #8511 as they stood before their stream-major mHC norm and 8-way head-EP commits. The per-change A/B added one commit of this PR at a time to that stack, and the byte-identity runs added this PR with the other changes named there. The device test suite ran on the integrated tree, which on top of the byte-identity tree carries the stream-major mHC norm operand (now in #8510, merged), the 8-way head-EP exchange (now in #8511), the one-kernel mHC Sinkhorn (#8594) and the Ulysses send buffer (#8595). On main every MAGI-2 file equals c1e84ce plus the five merged PRs, and the added lines of this branch are identical to the integrated tree's (including the 32/64 K tile of the test reference) except for the test fix, which was validated on this branch (see Test Result).
Test Result
test_native_compile*.pywithNameError: name 'device' is not defined, and they fail the same way on plain c1e84ce in this image; fix(device): register the device factory as the FX device builtin MooreThreads/torchada#124 (merged) fixes theNameError; the distributed test then stops at torch_musa, which does not support FSDP2 on a CPU mesh.cudacases of the*_on_devicetests intest_moe_routing_parity.py(route metadata, token-major forward and blocks, router operand) cover the routing side on a CUDA runner; the shared-expert device tests are MUSA-only, so its CUDA parity has no test.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*.pyFXNameErrorfrom torchada'storch.devicereplacement (fix(device): register the device factory as the FX device builtin MooreThreads/torchada#124 (merged) fixes theNameError; the distributed test then stops at torch_musa, which does not support FSDP2 on a CPU mesh), 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 c1e84ce.test_musa_device_shared_experts_match_the_concatenated_activation_bitwise[compiled], at group layout (29, 0, 0). The test's fixture set Inductor's deterministic flag once, but Dynamo restorestorch.use_deterministic_algorithmsafter every traced frame, which also resets that flag, so later compiles could pad the unaligned BF16 GEMMs by benchmarking and lower the two forms differently. On MTT S5000, forcing that padding makes the old test fail at (29, 0, 0) and forbidding it makes it pass; in 5 natural runs the GEMM was padded 3 times and the test failed exactly those 3 times, though one further pytest run padded it and passed, so the attribution is strong but not complete. With the test fix (0feba33, on c1e84ce), which passesdeterministic=Truein the compile options, the Test Plan gives 109 passed, 26 skipped, 0 failed, and the forced-padding and natural runs all pass.test_moe_routing_parity.pypassed with K=32 on both sides there and, in the integrated tree, with the 32/64 tile that the rebased reference now uses.benchmark_magi2_bf16_moe.py --mode routed: status pass at 2 and 4096 tokens.AI assistance: Claude Code drafted the change, the tests and this description and ran the MUSA validation listed above. I reviewed every change in this PR and re-ran the validation listed above.