Repository navigation
fix(mla): pass merge_attn_states' prefill_tokens_with_context at runtime - #2378
Merged
Merged
Conversation
merge_attn_states_kernel declared prefill_tokens_with_context as tl.constexpr.
Every MLA caller leaves it at its num_tokens default, so the kernel compiled
once per distinct batch token count -- a per-batch JIT on the serving path
that stalls every TP rank.
Kimi-K3 AgentX on MI355X, TP8, fresh triton cache: 274 of 291 in-window
triton compiles at c1 and 711 of 712 at c14 were this kernel, and requests
that overlapped one were ~0.8-0.9 s slower to first token than the same
request without it.
It is only compared against token_idx, a branch that was already runtime
(token_idx is program_id), so it becomes a plain int with
do_not_specialize: the generated IR differs only in reading the bound from
an argument instead of a constant.
- 60 batch sizes: 120 compiled kernels before, 6 after (OUTPUT_LSE x the
num_tokens ==1 / %16 / other classes)
- bitwise-identical outputs vs the old kernel over 2304 cases (bf16 and
fp8 e4m3fn output, +-inf LSE edges, prefill_tokens_with_context in
{None, 0, 1, T//2, T})
- tests/test_merge_attn_states.py: 22 passed
- do_bench, interleaved: within -1%..+4%, <= 1.4 us/call
- serving c1: 0 in-window compiles of this kernel (was 274/h), TTFT p50
back to the no-JIT level
Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Contributor
🏷️ CI GuideRuns automatically on every eligible PR before approval:
Heavy model tests:
|
valarLip
approved these changes
Sep 24, 2026
zhuyuhua-v
pushed a commit
that referenced
this pull request
Sep 24, 2026
…16, capture from bs=1 (#2382) * (recipe) Kimi-K3 AgentX: FP8 prefill attention, PrefillDelayer from C16, capture from bs=1 Adopt the configuration of the 2026-09-23 AgentX sweep (C1/4/14/16/48/56/72 on 8x MI355X): - ATOM_USE_FLYDSL_FP8_PREFILL_ATTN=1 on every band - ATOM_PREFILL_DECODE_INTERVAL=4 and ATOM_PREFILL_DELAYER_MAX_QUEUE_MS=5000 on CONC >= 16 - CUDA-graph capture sizes start at 1 (seq -s, 1) - container prerequisite: the 92 FlyDSL FP8 prefill kernels are AOT-built (ROCm/aiter#5796) and ATOM carries the merge_attn_states runtime-arg fix (#2378), with a one-line check - ReplaySSM paragraph now lists C14 among the on-bands, matching the table Against the current recipe's best runs: C48-C72 +10-15% total and +10-17% output throughput, ITL -14-18%; C1/C4 ITL p90 -12-13%; C14 TTFT p90 -19%. The delayer raises TTFT from C16 up (p50 +37-92%). The launcher's exported env and argv match the sweep's server script on all 14 bands, apart from default-equal or logging-only values. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com> * (recipe) Kimi-K3 AgentX: keep the change to env vars and launch args Drop the FP8 prefill precompile prerequisite subsection and restore the section-0 intro. Revert the ReplaySSM paragraph edit. The recipe diff is now limited to ATOM_USE_FLYDSL_FP8_PREFILL_ATTN=1, the C>=16 PrefillDelayer variables and capture sizes starting at 1, plus their table rows. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com> --------- Co-authored-by: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
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.
What
merge_attn_states_kerneldeclaredprefill_tokens_with_contextastl.constexpr. Every MLA caller (attention_mla.py, the three chunked-prefill merges) leaves it at itsnum_tokensdefault, so the kernel was compiled once per distinct batch token count: a JIT compile on the serving path, which stalls every TP rank.The value is only compared against
token_idx(tl.program_id), a branch that was already a runtime branch. This PR makes it a plain int withdo_not_specialize, so the generated IR differs only in reading the bound from a kernel argument instead of a constant.Why it matters
Kimi-K3 AgentX (aiperf, 3600 s window) on MI355X, TP8, with a fresh triton cache:
merge_attn_states_kernelRequests whose prefill overlapped one of these compiles were ~0.8–0.9 s slower to first token than the same request (paired by conversation/turn) without one. At c14, 445 of 520 requests were hit and TTFT p50 doubled.
Because the key is the token count, a warm cache never fully covers it: the observed values ran from 2 to 8192.
Validation
OUTPUT_LSEon/off: 120 compiled kernels before, 6 after (OUTPUT_LSE× thenum_tokens==1 / %16 / other classes).e4m3fn) output, ±inf LSE edge rows,prefill_tokens_with_context∈ {None, 0, 1, T//2, T}.tests/test_merge_attn_states.py: 22 passed.do_bench, interleaved A/B × 5 rounds, H=12 D=128:OUTPUT_LSE=Falsewithin +0–2%,Truewithin −1%..+4% (≤ 1.4 µs/call).The vLLM plugin paths call vLLM's own
merge_attn_states, where this argument is already a runtime int, so they are unaffected.