Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 6 additions & 3 deletions tensorrt_llm/_torch/modules/ATTENTION_DEVELOPER_GUIDE.md
Original file line number Diff line number Diff line change
Expand Up @@ -337,9 +337,12 @@ cache (override with `TLLM_K3_MLA_GEN_BACKEND=trtllm-gen`; other values are
rejected at model build). FP8 KV cache forces `trtllm-gen`. K3 also installs a
per-batch policy that falls back to `trtllm-gen` for mixed
context/generation batches and multi-token generation (speculative
verification), keeping `cute-dsl` for plain one-token-per-request decode. A
mixed H=96 batch remains on `cute-dsl`: TRTLLM-Gen may select a 64-head Q tile,
which does not divide 96 after K3's head padding removal.
verification), keeping `cute-dsl` for plain one-token-per-request decode.
Any H=96 batch (K3's attention-DP shape) remains on `cute-dsl` regardless of
batch composition: TRTLLM-Gen may select a 64-head Q tile, which does not
divide 96 after K3's head padding removal, and its decode gate rejects
`64 < num_heads_q < 128` — so falling back there would fail engine
initialization (this covers attention-DP speculative verification).

The FMHA package is split by role:

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -77,19 +77,21 @@ def _kimi_k3_mla_decode_backend_policy(
CuTe-DSL reuses one staged page table across MLA layers for a
generation-only, one-token-per-request batch. Other mixed batches repeat
the staging copies in every MLA layer and regress time to first token, so
they fall back to TRTLLM-Gen. The H=96 path is the correctness exception:
TRTLLM-Gen may select a 64-head Q tile, which does not divide 96 and
produces an invalid configuration after K3's head padding was removed.

The CuTe-DSL kernel itself accepts multi-token queries, but K3's decode
tuning covers only the one-token-per-request regime, so generation-only
speculative verification also falls back.
they fall back to TRTLLM-Gen. The H=96 path is the correctness exception
and applies to EVERY fallback candidate: TRTLLM-Gen may select a 64-head
Q tile, which does not divide 96 (invalid after K3's head padding was
removed), and its decode gate rejects 64 < num_heads_q < 128 outright —
falling back would fail engine initialization. H=96 per rank is K3's
attention-DP shape, so this keeps attention-DP + speculative
verification (a generation-only multi-token batch) on CuTe-DSL, which
accepts multi-token queries; K3's decode tuning preference for
TRTLLM-Gen only applies where TRTLLM-Gen is valid at all.
"""
is_single_token_generation = num_gen_tokens == metadata.num_generations
requires_cute_dsl_for_mixed_batch = metadata.num_contexts > 0 and num_heads == 96
requires_cute_dsl = num_heads == 96
if (
requested_backend == "cute-dsl"
and not requires_cute_dsl_for_mixed_batch
and not requires_cute_dsl
and (metadata.num_contexts > 0 or not is_single_token_generation)
):
return "trtllm-gen"
Expand Down
9 changes: 7 additions & 2 deletions tests/unittest/_torch/modules/test_kimi_k3_mla_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,12 @@ def test_select_kimi_k3_mla_generation_backend_uses_trtllm_gen_for_fp8_kv_cache(
("cute-dsl", 0, 4, 4, 96, "cute-dsl"),
("cute-dsl", 1, 3, 3, 12, "trtllm-gen"),
("cute-dsl", 1, 3, 3, 96, "cute-dsl"),
("cute-dsl", 0, 4, 8, 96, "trtllm-gen"),
# Generation-only multi-token (speculative verification): H=96 must
# stay on cute-dsl — trtllm-gen's decode gate rejects that head
# count, so falling back fails engine init. Smaller per-rank head
# counts (non-attention-DP shapes) keep the tuning fallback.
("cute-dsl", 0, 4, 8, 96, "cute-dsl"),
("cute-dsl", 0, 4, 8, 12, "trtllm-gen"),
("trtllm-gen", 1, 3, 3, 96, "trtllm-gen"),
],
)
Expand All @@ -78,7 +83,7 @@ def test_kimi_k3_mla_decode_backend_policy_by_batch_shape(
num_heads: int,
expected_backend: str,
) -> None:
"""K3 falls back outside plain decode except for unsafe H=96 mixed batches."""
"""K3 falls back outside plain decode except when H=96 breaks trtllm-gen."""
assert (
_kimi_k3_mla_decode_backend_policy(
requested_backend,
Expand Down
Loading