Skip to content

[Perf][DSv4.1] Fold the mHC post block into the delayed pre projection - #56633

Merged
zyongye merged 8 commits into
vllm-project:mainfrom
zyongye:perf/dsv41-fused-post-pregemm
Sep 14, 2026
Merged

zyongye merged 8 commits into
vllm-project:mainfrom
zyongye:perf/dsv41-fused-post-pregemm

Conversation

@zyongye

@zyongye zyongye commented Sep 12, 2026 •

Copy link
Copy Markdown
Member

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_tilelang computes 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_tilelang runs 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's mhc_post_tilelang + split-k GEMM, so only the small-batch path changes shape.

Both V4.1 seams route through it, except one:

  • Attention seam — fuses, unless Engram is injected. Engram reads the residual between the post and the pre, so the post cannot move into the projection there; that branch keeps the separate kernels.
  • FFN seam — always fuses.
  • DSpark draft reuses 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_keys takes extra_splits and 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 at hc_mult=4, CUDA-graph device time:

tokens 1 2 4 8 16 24 32 48 64
inherited config, 5120 8.21 8.01 8.42 8.82 10.25 12.73 13.51 16.81 19.70
tuned, 5120 8.02 8.00 8.83 8.84 9.85 10.47 11.47 fallback fallback
unfused, 5120 12.74 12.52 12.72 12.92 13.32 13.32 13.37 13.53 13.49

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_n 2 below 16 and 6 through 32, n_thr 128), within ~1-2% of the per-shape optimum for one extra compiled variant. Same bands hold at hidden_size 7168.

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:

tokens 1 2 4 8 16 24 32 48
fused 8.44 8.23 8.83 9.03 9.85 10.61 11.29 13.55
unfused, split as configured today 13.13 12.74 12.76 13.33 13.35 13.55 13.56 13.35
unfused, split=64 11.09 11.09 11.29 11.66 11.50 11.69 11.62 11.72
fused vs tuned baseline 1.31x 1.35x 1.28x 1.29x 1.17x 1.10x 1.03x 0.86x

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

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_mhc lazily 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).

tokens 1 2 4 8 16 24 32 48 64 128
unfused 13.13 12.84 13.02 13.54 13.32 13.54 13.54 13.55 13.56 13.97
fused TileLang (this PR) 8.43 8.22 8.83 8.84 9.65 10.47 11.28 13.10 13.54 13.96
Mega-mHC (#56255) 10.19 10.46 10.07 10.67 10.47 10.47 10.88 11.09 11.26 11.82

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.profiler kernel-time breakdown agrees (at 1 token: 4.32 µs fused kernel + 5.03 µs epilogue, against a single 10.44 µs sm100_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

.venv/bin/python -m pytest tests/kernels/test_mhc_kernels.py -v

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 to mhc_post_tilelang + mhc_pre_delayed_tilelang, and the coefficients match an FP32-post torch reference to 1e-5 on the fused path. Plus a torch.library.opcheck case and a CUDA-graph capture/replay case, since decode replays this op from a graph. test_deepseek_v41_decoder_mixes_match_torch now covers both routes through the decoder.

Test Result

164 passed, 23 skipped in 53.80s     # tests/kernels/test_mhc_kernels.py, GB300

pre-commit run --all-files clean 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: 0 and 0 null answers in every run:

arm runs exact_match/flexible
this PR 3 0.9242, 0.9325, 0.9204
unfused baseline, same serve config 2 0.9348, 0.9249

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_tilelang the model ran per aux-capture layer is recomputing a residual that already exists. The second commit takes mean(dim=1) from the fused call behind a capture_previous_aux flag 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_buffer copy when aux hidden states are present: that depends on its vllm/v1/worker/gpu/model_runner.py edit, 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 a write_aux specialization 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:

tokens 1 8 32 128 1024 4096
seam + separate mean 10.55 11.50 13.94 17.65 60.25 209.32
seam + folded mean 8.43 9.04 11.50 13.74 40.88 140.63
saved (µs) 2.13 2.46 2.45 3.90 19.38 68.70

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.mean on 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 (not synthetic), 8 fixed prompts at temperature 0, same branch with and without the commit:

drafts draft tokens accepted rate length per position
neither commit 106 318 204 0.6415 2.9245 92 / 64 / 48
harvest 106 318 204 0.6415 2.9245 92 / 64 / 48
harvest + folded mean 106 318 204 0.6415 2.9245 92 / 64 / 48

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_aux pins the same thing at the unit level with atol=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_tokens 57344, concurrency 128, same TP4 GB300 node:

arm headline answered-only pass@4 request errors truncated null answers empty@stop
DSpark k=5, real verification 0.9104 0.9138 0.9545 0 3 3 0
no spec decode 0.9129 0.9163 0.9495 0 3 3 0

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

@claude claude 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.

Claude Code Review

This pull request is from a fork — automated review is disabled. A repository maintainer can comment @claude review to run a one-time review.

@mergify mergify Bot added deepseek Related to DeepSeek models DSv4.1 Related to DeepSeek-V4.1 models labels Sep 12, 2026
@zyongye

zyongye commented Sep 13, 2026

Copy link
Copy Markdown
Member Author

/ci run

@github-actions

Copy link
Copy Markdown

✅ Triggered Buildkite CI #88578 for commit e500862b3a1e.

@zyongye
zyongye force-pushed the perf/dsv41-fused-post-pregemm branch from e500862 to c063aa8 Compare September 13, 2026 01:02
@zyongye

zyongye commented Sep 13, 2026

Copy link
Copy Markdown
Member Author

/ci run

@github-actions

Copy link
Copy Markdown

✅ Triggered Buildkite CI #88604 for commit f6baeecafbff.

@zyongye
zyongye force-pushed the perf/dsv41-fused-post-pregemm branch from f6baeec to 3432199 Compare September 13, 2026 22:14
@zyongye

zyongye commented Sep 14, 2026

Copy link
Copy Markdown
Member Author

/ci run

@zyongye
zyongye enabled auto-merge (squash) September 14, 2026 02:14
@github-actions github-actions Bot added the ready ONLY add when PR is ready to merge/full CI is needed label Sep 14, 2026
@github-actions

Copy link
Copy Markdown

✅ Triggered Buildkite CI #88704 for commit 3432199c6dbe.

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
zyongye force-pushed the perf/dsv41-fused-post-pregemm branch from 3432199 to 9bf278d Compare September 14, 2026 06:32
@zyongye

zyongye commented Sep 14, 2026

Copy link
Copy Markdown
Member Author

/ci run

@github-actions

Copy link
Copy Markdown

✅ Triggered Buildkite CI #88749 for commit 9bf278d7b256.

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>
@zyongye

zyongye commented Sep 14, 2026

Copy link
Copy Markdown
Member Author

/ci run

@github-actions

Copy link
Copy Markdown

✅ Triggered Buildkite CI #88880 for commit 0f92779c48d4.

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

Labels

deepseek Related to DeepSeek models DSv4.1 Related to DeepSeek-V4.1 models ready ONLY add when PR is ready to merge/full CI is needed

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants