Repository navigation
[Perf][DSv4.1] Fold the mHC post block into the delayed pre projection - #56633
Merged
zyongye merged 8 commits intoSep 14, 2026
Merged
Conversation
zyongye
requested review from
AndreasKaratzas,
WoosukKwon,
mgoin,
tlrmchlsmth and
yewentao256
as code owners
September 12, 2026 19:28
Member
Author
|
/ci run |
|
✅ Triggered Buildkite CI #88578 for commit |
ywang96
approved these changes
Sep 13, 2026
zyongye
force-pushed
the
perf/dsv41-fused-post-pregemm
branch
from
September 13, 2026 01:02
e500862 to
c063aa8
Compare
Member
Author
|
/ci run |
|
✅ Triggered Buildkite CI #88604 for commit |
zyongye
force-pushed
the
perf/dsv41-fused-post-pregemm
branch
from
September 13, 2026 22:14
f6baeec to
3432199
Compare
Member
Author
|
/ci run |
zyongye
enabled auto-merge (squash)
September 14, 2026 02:14
|
✅ Triggered Buildkite CI #88704 for commit |
The tile, split-k and block size carried a TODO to autotune them and were
duplicated between MhcFusedTileLangKernel.dispatch, its launch spec and
the wrapper's token cutoff, so the three could drift. Sweeping all of them
on GB300 at hc_mult 4 against the separate post + split-k GEMM they
replace, at hidden_size 5120 and 7168:
- one fixed tile starves the projection as the batch grows: the
inherited config is 18% off the per-shape optimum at 24-32 tokens
- a 128-thread block with an 8-way split wins from 1 token upward
- the 16-token cutoff was measured against a split-k GEMM using the
bucketed split count from compute_mhc_pre_num_splits, which is itself
1.14-1.19x off below 128 tokens. Against a GEMM given the split its
own estimator computes, the fused path is 1.3x ahead at 1-8 tokens,
1.17x at 16, break-even at 32 and 0.86x by 48
So the bands are tile_n 2 below 16 tokens and 6 through 32, at 128
threads, and the cutoff moves 16 -> 32. mhc_fused_post_pre_split_config is
now the one source of truth that the compile key, the launch and the
wrapper all read, and it declines shapes the kernel cannot tile evenly
(mix_size % tile_n, hidden_size % (n_splits * n_thr)) instead of silently
dropping that work.
Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Signed-off-by: Yongye Zhu <yongye@inferact.ai>
Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
DeepSeek-V4.1 runs an mHC post block and then the next sublayer's delayed pre block at every seam, so each seam pays for a post kernel, a split-k projection that re-reads the residual streams the post kernel just wrote, and the pre epilogue. DSv4 already avoids that middle read through mhc_fused_tilelang; V4.1 could not reuse it because its pre carries the previous sublayer's pre-mix in and hands the next one back. mhc_fused_post_pre_delayed_tilelang runs that fused kernel and then the delayed pre epilogue, and both V4.1 seams route through it. Above the fused kernel's token range it falls back to the same post kernel and split-k GEMM as today, so only the small-batch path changes shape. The Engram seam keeps the separate kernels, since Engram reads the residual between the post and the pre. The delayed epilogue is compiled per bucketed split count, so the fallback projection uses compute_mhc_pre_num_splits rather than the raw estimate the V4 path uses, and the decoder registers the fused GEMM's own split-k factors for warmup -- without that the first decode step JITs inside the capture warmup. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Signed-off-by: Yongye Zhu <yongye@inferact.ai> Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
Every layer's post now runs inside the next layer's fused pre, so the standalone mhc_post_tilelang the model ran per aux-capture layer was recomputing a residual that already existed. Take the mean over the hc streams from the fused call instead, behind a capture_previous_aux flag, and drop the duplicate post. The Engram seam captures before the injection, matching the standalone post it replaces: aux consumers read the stream before Engram touches it. The last layer on a rank has no successor to fold its post into, so it keeps the final standalone post and takes its mean from there. Ports the approach of vllm-project#55575 (DeepSeek-V4) to V4.1, which becomes possible here because this branch gives V4.1 the same fused seam. DSpark on DeepSeek-V4.1-Flash, TP4 GB300, k=3 probabilistic drafting with real verification, 8 fixed prompts at temperature 0: byte-identical output and identical acceptance against the same branch without this commit -- 106 drafts, 318 draft tokens, 204 accepted (0.6415, length 2.92), per position 92/64/48. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Signed-off-by: Yongye Zhu <yongye@inferact.ai> Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
The collapse in the pre epilogue already reads every post-mapped residual stream to weight them by the carried pre-mix. The aux hidden state draft models consume is the unweighted mean of those same streams, so accumulate it in the same pass instead of running mean(dim=1) over the residual again. Gated by a write_aux specialization so non-draft setups compile and run exactly what they did before, and the buffer is only allocated when the caller asks for it. Bit-exact against torch's mean at 1/8/33/128/1024 tokens on both 5120 and 7168 hidden sizes, and capturing leaves every other output unchanged. Saves the whole cost of the separate reduction: 2.1-2.5 us per capture seam at decode widths, 19 us at 1024 tokens and 69 us at 4096. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Signed-off-by: Yongye Zhu <yongye@inferact.ai> Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
The previous commit folded the stream mean into the collapse only in the epilogue that also applies RMSNorm, so a caller without a norm weight fell back to a torch mean plus a copy into the output buffer. Give the unnormalized epilogue the same write_aux path and drop the fallback: both now accumulate the mean from the fragment the collapse already loaded. Bit-exact against torch's mean on both paths at 1/8/33/128/1024 tokens and 5120/7168 hidden, with every other output unchanged when capturing. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Signed-off-by: Yongye Zhu <yongye@inferact.ai> Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
vllm/model_executor/kernels/mhc/__init__.py re-exports tilelang.py, and tilelang.py imports tilelang_kernels only inside the functions that launch kernels -- so importing the mHC package does not require tilelang, which vllm/tilelang_utils raises on when it is missing under CUDA. Reaching for the shared split config at module scope broke that: the package started pulling in tilelang and applying @tilelang_jit to every kernel at import, 1.36 s of it, and hard-failing wherever tilelang is absent. Move both imports back inside their callers, matching what the DSv4 model already does with the same module. Importing the package is 6 ms again and neither tilelang nor tilelang_kernels lands in sys.modules. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Signed-off-by: Yongye Zhu <yongye@inferact.ai> Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
…test vllm-project#56741 renamed vllm/models/deepseek_v4_1 to deepseek_v41. The monkeypatch target for the fused seam is a string, so neither that rename nor this rebase could follow it, and the decoder test imported a module that no longer exists. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com> Signed-off-by: Yongye Zhu <yongye@inferact.ai> Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
zyongye
force-pushed
the
perf/dsv41-fused-post-pregemm
branch
from
September 14, 2026 06:32
3432199 to
9bf278d
Compare
Member
Author
|
/ci run |
|
✅ Triggered Buildkite CI #88749 for commit |
The stub's __call__ took *args only and returned five values, so reading aux hidden states out of the fused post broke it: capture_previous_aux is keyword-only and cannot land in *args. Accept **kwargs and return the sixth value, which keeps the stub tolerant of the next keyword-only flag rather than pinned to today's signature. Engram hashes and the mask are still the last positional arguments, so args[-2:] is unchanged. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com> Signed-off-by: Yongye Zhu <yongye@inferact.ai> Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
Member
Author
|
/ci run |
|
✅ Triggered Buildkite CI #88880 for commit |
ItsRoy69
pushed a commit
to ItsRoy69/vllm
that referenced
this pull request
Sep 15, 2026
vllm-project#56633) Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
keneoneth
pushed a commit
to keneoneth/vllm
that referenced
this pull request
Sep 16, 2026
vllm-project#56633) Signed-off-by: Yongye Zhu <zyy1102000@gmail.com>
This was referenced Sep 23, 2026
This was referenced Oct 9, 2026
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
DeepSeek-V4.1 runs an mHC post block and then the next sublayer's delayed pre block at every seam. Each seam therefore pays for three things: a post kernel, a split-k projection that re-reads the residual streams the post kernel just wrote, and the pre epilogue.
DSv4 already avoids the middle read:
mhc_fused_tilelangcomputes the updated residual in registers and feeds the pre-norm GEMM from there. V4.1 could not reuse it, for the same reason it needed its own pre in the first place — it collapses each sublayer's input with the previous sublayer's pre-mix and hands the next one back, and the DSv4 fused wrapper neither accepts an incoming pre-mix nor returns one.Approach
mhc_fused_post_pre_delayed_tilelangruns the existing fused kernel and then the delayed pre epilogue (MHC_PRE_NORM_KERNEL,use_pre_mix_in/save_pre_mix), returning(residual, post_mix, comb_mix, layer_input, next_pre_mix). Above the fused kernel's token range it falls back to exactly today'smhc_post_tilelang+ split-k GEMM, so only the small-batch path changes shape.Both V4.1 seams route through it, except one:
DeepseekV4DecoderLayer, so it inherits both.The pre epilogue is JIT-compiled per split-k factor and warmed up from a token sweep that never produces the factors this GEMM picks, so
MHCPreNormKernel.get_warmup_keystakesextra_splitsand the decoder registers them. Without it the first decode step JITs inside the capture warmup.Tuning
The tile, split-k and block size were inherited from DSv4 unmeasured (the code carried a
TODO(gnovack): investigate autotuning these heuristics). Sweeping all three against the separate post + split-k GEMM, on GB300 athc_mult=4, CUDA-graph device time:Two things were wrong with the inherited constants: the 16-token cutoff was early (the fused path still wins 1.2x at 32), and one fixed tile starves the projection of parallelism as the batch grows — it is 18% off at 24-32 tokens and slower than the unfused path past 48. The tuned heuristic is two bands (
tile_n2 below 16 and 6 through 32,n_thr128), within ~1-2% of the per-shape optimum for one extra compiled variant. Same bands hold athidden_size7168.The cutoff is 32, and it is set against a corrected baseline. vLLM configures the TF32 projection's split count through
compute_mhc_pre_num_splits, which buckets a good estimate into {1, 4, 16} to bound startup specializations — and that bucketing is itself 1.14-1.19x off the optimum below 128 tokens (measured: the estimator says 80 splits at one token, the bucket gives 16, the optimum is 64; rounding the estimate down to a power of two reproduces the measured best at every shape I tried). Against a GEMM given the split its own estimator computes:So the honest claim is 1.3x at decode widths, 1.17x at 16, break-even at 32. The cutoff sits at the break-even point rather than where a slow baseline would have put it (48, which is a 14% loss against a tuned GEMM). Fixing that split heuristic is worth 5-19% on every unfused seam — all prefill, all large-batch decode, and DSv4 as well as DSv4.1 — but it belongs in its own PR, not bundled here, since it changes a shared helper and a second model's hot path.
It is shared with the DSv4 wrapper, which uses the same kernel at the same shapes. The heuristic now also declines shapes the kernel cannot tile evenly (
mix_size % tile_n,hidden_size % (n_splits * n_thr)) and falls back instead — previously those silently dropped whole tiles and k-slices.Relationship to the other open mHC PRs
mhc_shifted_post_predispatches to DeepGEMMmega_mhcwhendevice_capability_family(100) and hidden_size % 1024 == 0 and hc_mult == 4, and falls back to the unfused pair otherwise — which is what this PR replaces. The two are complementary and land cleanly together: benchmarked head to head below, the TileLang fusion is faster at small batch, Mega-mHC is faster at large batch. If both land, the right end state is one dispatcher picking between them by token count, and I am happy to write that as a follow-up (or fold this into [DSv4.1] Integrate Mega-mHC from DeepGEMM #56255, author's preference).mhc_fused_post_pre_gemm_sqrsum, on_aiter_ops.py/mhc/aiter.py/layers/mhc.py/amd/model.py. No file overlap with this PR beyond the shared test file.Mega-mHC comparison
DeepGEMM 2.8.0 (
66081d4c), GB300,hc_mult=4,hidden_size=5120, CUDA-graph device time per seam. Both paths are graph-capturable;mega_mhclazily allocates per-stream split barriers and asserts that happens outside capture, so it must be warmed on the capture stream first (which vLLM's pre-capture warmup iteration does).This PR is 1.08–1.27x faster than Mega-mHC at 1–16 tokens; they cross at ~24; Mega-mHC is 1.18–1.20x faster at 48+. A
torch.profilerkernel-time breakdown agrees (at 1 token: 4.32 µs fused kernel + 5.03 µs epilogue, against a single 10.44 µssm100_mega_mhc_impl).Numerically the two differ in one place: the fused TileLang kernel projects the FP32 post result it holds in registers, while Mega-mHC and the unfused path project its BF16 rounding. Against an FP32 reference at 8 tokens, max abs error: residual 9.77e-4 for all three; coefficients 9.25e-4 / 1.78e-3 (this PR) against 1.98e-4 / 3.46e-4 (Mega-mHC and unfused). This is inherited from the DSv4 kernel, not new — the stored residual streams are bit-identical either way.
Test Plan
New
test_deepseek_v41_mhc_fused_post_pre_delayed(0/1/8/128 tokens x 4096/7168 hidden x carried/absent pre-mix) asserts the post-mapped residual and the collapsed layer input are bit-identical tomhc_post_tilelang+mhc_pre_delayed_tilelang, and the coefficients match an FP32-post torch reference to 1e-5 on the fused path. Plus atorch.library.opcheckcase and a CUDA-graph capture/replay case, since decode replays this op from a graph.test_deepseek_v41_decoder_mixes_match_torchnow covers both routes through the decoder.Test Result
pre-commit run --all-filesclean on the changed files (ruff, ruff-format, mypy-3.10).Model evaluation
DeepSeek-V4.1-Flash, TP4 GB300, gsm8k 5-shot via
mqa-run --api-type completions --temperature 0 --max-tokens 1024, n=1319,num_request_errors: 0and 0 null answers in every run:The baseline arm's own two runs differ by 0.99 pt, so the arms' ranges overlap and the 0.42 pt difference in means sits inside the measured noise floor for this config (batching nondeterminism at concurrency 32). Dashboard history for this model is median 0.9314, p25–p75 0.9287–0.9333.
Second and third commits: aux hidden states for draft models
Once the post lives inside the next layer's pre, the standalone
mhc_post_tilelangthe model ran per aux-capture layer is recomputing a residual that already exists. The second commit takesmean(dim=1)from the fused call behind acapture_previous_auxflag and drops the duplicate post — the approach of #55575 (DeepSeek-V4), which only becomes available to V4.1 once V4.1 has the fused seam.Two seams keep the standalone post: the Engram seam captures before the injection (matching the post it replaces — aux consumers read the stream before Engram touches it), and the last layer on a rank has no successor to fold into, so it takes its mean from the final post that runs anyway.
I did not port #55575's companion change that skips the
_mtp_hidden_buffercopy when aux hidden states are present: that depends on itsvllm/v1/worker/gpu/model_runner.pyedit, which is model-agnostic and would collide with #55575 in a shared file. It should follow from that PR, not this one.Third commit: fold the stream mean into the collapse
That leaves one pass still standing:
mean(dim=1)over[T, hc, H]. The collapse in the pre epilogue already reads every one of those streams to weight them by the carried pre-mix, and the aux value is the unweighted mean of the same data — so the third commit accumulates it in that pass, behind awrite_auxspecialization so non-draft setups compile and run exactly what they did before. The buffer is allocated only when the caller asks.Standalone
mean(dim=1)is not cheap at scale (measured on GB300,[T, 4, 5120]bf16): 2.10 µs at 1 token, 17.9 µs at 1024, 72.1 µs at 4096, 141.8 µs at 8192 — per capture layer, per step. Folded, the aux-capture seam costs:Both epilogues carry the path — the one that applies RMSNorm and the plain one — so no caller falls back to a torch mean. Bit-exact against
torch.meanon both, at 1/8/33/128/1024 tokens and 5120/7168 hidden, with every other output of the op unchanged when capturing is on. The Engram seam still uses a separate mean, since its aux is the pre-injection post and the epilogue runs after Engram.DSpark verification
DeepSeek-V4.1-Flash, TP4 GB300,
{"method":"dspark","num_speculative_tokens":3,"draft_sample_method":"probabilistic"}with real verification (notsynthetic), 8 fixed prompts attemperature 0, same branch with and without the commit:Byte-identical generated text (sha256 of the concatenated completions matches) and identical acceptance counters, which is the end-to-end statement that the harvested tensor is the one the draft model used to get.
test_deepseek_v41_capture_previous_auxpins the same thing at the unit level withatol=0, for both the fused and Engram seams.That probe pins the tensor exactly but is only 8 prompts, so both arms were also run through GPQA-Diamond at full size — 198 questions x 4 epochs = 792 samples per arm, chat path with thinking on, temperature 1.0, top_p 0.95,
max_tokens57344, concurrency 128, same TP4 GB300 node:Acceptance over that run, from the engine's own counters rather than a synthetic profile: 1,467,764 drafts, 7,338,820 draft tokens, 3,207,521 accepted — rate 0.4371, acceptance length 3.185 of a possible 6.
Both arms are this branch, so this is not a before/after for the PR; the before/after is the gsm8k table above. What it establishes is that the code these two commits touch carries 1.47M drafts at an accuracy indistinguishable from the no-draft arm: 2 questions of 792, where one question is 0.126 pt and the sign flips between exact_match and pass@4. Finished generations ran to a median of 2239 tokens with a p90 of 16232 and a max of 56320 against the 57344 cap, so the three truncations per arm are a genuine budget limit rather than a loop — identical in both arms, so they do not move the comparison.
AI assistance
This change was written with AI assistance (Claude Code). I reviewed every changed line, ran the tests and benchmarks above myself, and can defend the change end to end.
🤖 Generated with Claude Code