Skip to content

fix(mla): pass merge_attn_states' prefill_tokens_with_context at runtime - #2378

Merged
valarLip merged 1 commit into
mainfrom
fix/merge-attn-states-runtime-ptwc
Sep 24, 2026
Merged

valarLip merged 1 commit into
mainfrom
fix/merge-attn-states-runtime-ptwc

Conversation

@gbyu-amd

@gbyu-amd gbyu-amd commented Sep 23, 2026 •

Copy link
Copy Markdown
Contributor

What

merge_attn_states_kernel declared prefill_tokens_with_context as tl.constexpr. Every MLA caller (attention_mla.py, the three chunked-prefill merges) leaves it at its num_tokens default, 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 with do_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:

c1 c14
in-window triton compiles 291 / h 712 in the first ~20 min
of which merge_attn_states_kernel 274 711

Requests 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

  • Compile count, 60 random batch sizes × OUTPUT_LSE on/off: 120 compiled kernels before, 6 after (OUTPUT_LSE × the num_tokens ==1 / %16 / other classes).
  • Bitwise identical to the old kernel over 2304 cases: T ∈ {1..4099}, H ∈ {1,5,12,16}, D ∈ {128,192}, bf16 and fp8 (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=False within +0–2%, True within −1%..+4% (≤ 1.4 µs/call).
  • Serving, same c1 workload: 0 in-window compiles of this kernel (was 274/h). TTFT p50 is back at the no-JIT level: 735 ms, vs 1418 ms before on the same requests.

The vLLM plugin paths call vLLM's own merge_attn_states, where this argument is already a runtime int, so they are unaffected.

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>
@github-actions

Copy link
Copy Markdown
Contributor

🏷️ CI Guide

Runs automatically on every eligible PR before approval:

  • ✅ Pre Checkin: Black, Ruff, catalog schema validation, non-GPU unit tests

Heavy model tests:

  • ✅ Run after the PR is approved and Pre Checkin passes
  • ✅ Run immediately when an approval review is submitted
  • ✅ Can be requested before approval with labels
Label Tests
ci:full Run all heavy PR model tests: native ATOM, vLLM, and SGLang
ci:atom Run native ATOM model accuracy tests
ci:vllm Run ATOM vLLM OOT model accuracy tests
ci:sglang Run ATOM SGLang model accuracy tests

Heavy jobs are skipped when the PR is not approved and no matching ci:* label is present.
Add labels via the sidebar or gh pr edit 2378 --add-label <label>

@valarLip
valarLip merged commit acda72f into main Sep 24, 2026
72 of 75 checks passed
@valarLip
valarLip deleted the fix/merge-attn-states-runtime-ptwc branch September 24, 2026 03:33
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>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants