From fa1620da89c2964878cd8f1f27e874ad7b68559e Mon Sep 17 00:00:00 2001 From: Michal Guzek Date: Tue, 25 Aug 2026 22:41:44 -0700 Subject: [PATCH] [None][fix] Keep CuTe-DSL MLA decode for Kimi K3 H=96 speculative-verify batches The Kimi K3 MLA decode-backend policy kept a requested cute-dsl backend only for plain single-token decode batches and for mixed batches with 96 query heads; every other batch shape fell back to trtllm-gen. Generation-only multi-token batches (speculative verification, e.g. suffix-automaton drafting) at 96 heads -- K3's per-rank head count under attention-DP -- were therefore routed to trtllm-gen, whose decode gate rejects 64 < num_heads_q < 128, and engine initialization failed: ValueError: trtllm-gen MLA decode does not support 64 < num_heads_q < 128; got num_heads_q=96. Apply the H=96 exception to every fallback candidate instead of only mixed batches, so any batch shape trtllm-gen cannot serve stays on cute-dsl. Smaller per-rank head counts keep the existing fallback. Validated on a single B200 (SM100) with the SA harness (tests/integration/defs/kimi_k3_sa_harness.py) at TP=1 with attention-DP (96 heads per rank), a 4-layer truncated Kimi-K3, KIMI_K3_SPEC_MODE=sa and logits-parity checking: - without this change, the SA engine fails initialization with the ValueError above (the baseline single-token engine is unaffected); - with this change, the run passes spec-dec logits parity (4 prompts, 0 drift). The unit-test matrix in tests/unittest/_torch/modules/test_kimi_k3_mla_backend.py passes 12/12 with this change. Signed-off-by: Michal Guzek --- .../modules/ATTENTION_DEVELOPER_GUIDE.md | 9 ++++++--- .../kimi_k3_mla/kimi_k3_mla_attention.py | 20 ++++++++++--------- .../modules/test_kimi_k3_mla_backend.py | 9 +++++++-- 3 files changed, 24 insertions(+), 14 deletions(-) diff --git a/tensorrt_llm/_torch/modules/ATTENTION_DEVELOPER_GUIDE.md b/tensorrt_llm/_torch/modules/ATTENTION_DEVELOPER_GUIDE.md index 618ac27fd0e6..ae196276d5e2 100644 --- a/tensorrt_llm/_torch/modules/ATTENTION_DEVELOPER_GUIDE.md +++ b/tensorrt_llm/_torch/modules/ATTENTION_DEVELOPER_GUIDE.md @@ -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: diff --git a/tensorrt_llm/_torch/modules/kimi_k3_mla/kimi_k3_mla_attention.py b/tensorrt_llm/_torch/modules/kimi_k3_mla/kimi_k3_mla_attention.py index 889ac18c7efa..03bcc9004801 100644 --- a/tensorrt_llm/_torch/modules/kimi_k3_mla/kimi_k3_mla_attention.py +++ b/tensorrt_llm/_torch/modules/kimi_k3_mla/kimi_k3_mla_attention.py @@ -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" diff --git a/tests/unittest/_torch/modules/test_kimi_k3_mla_backend.py b/tests/unittest/_torch/modules/test_kimi_k3_mla_backend.py index 937daeabb2f9..67ec3ea13a4b 100644 --- a/tests/unittest/_torch/modules/test_kimi_k3_mla_backend.py +++ b/tests/unittest/_torch/modules/test_kimi_k3_mla_backend.py @@ -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"), ], ) @@ -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,