From 70aff584b8e865b47018f39e5e8f137b6db58626 Mon Sep 17 00:00:00 2001 From: William Zhang <133824995+2ez4bz@users.noreply.github.com> Date: Mon, 17 Aug 2026 22:59:41 -0700 Subject: [PATCH 1/2] [None][fix] Fix FlashInfer shared-KV speculative decode * Why? Overlap-scheduled speculative decoding with a shared KV cache can use draft positions beyond the logical sequence length. FlashInfer exposed only logical generation pages, and did not reflect scheduler offsets in its live KV metadata, while Gemma4 sized RoPE only to the logical limit. * What? Expose every reserved generation page for shared-KV overlap decoding while tracking logical KV lengths separately. Apply and restore overlap scheduler offsets to append positions and trtllm-gen decode lengths, and add speculative RoPE headroom to Gemma4 target and assistant models Signed-off-by: William Zhang <133824995+2ez4bz@users.noreply.github.com> --- .../_torch/attention_backend/flashinfer.py | 164 +++++++++++++++-- tensorrt_llm/_torch/metadata.py | 7 + tensorrt_llm/_torch/models/modeling_gemma4.py | 29 ++- .../_torch/models/modeling_speculative.py | 8 +- .../_torch/pyexecutor/model_engine.py | 43 ++++- .../test_llm_api_pytorch_multimodal.py | 3 - .../attention/test_flashinfer_attention.py | 169 +++++++++++++++++- .../_torch/modeling/test_modeling_gemma4.py | 38 +++- 8 files changed, 435 insertions(+), 26 deletions(-) diff --git a/tensorrt_llm/_torch/attention_backend/flashinfer.py b/tensorrt_llm/_torch/attention_backend/flashinfer.py index d49ab084a128..129c18c607c2 100644 --- a/tensorrt_llm/_torch/attention_backend/flashinfer.py +++ b/tensorrt_llm/_torch/attention_backend/flashinfer.py @@ -89,6 +89,21 @@ def _slice_paged_kv_cache_heads( return paged_kv_cache[tuple(index)] +def _get_page_table_num_blocks(kv_cache_manager, request_ids, + logical_num_blocks: list[int], + num_contexts: int) -> list[int]: + """Keep context rows logical and expose every reserved generation page.""" + generation_request_ids = request_ids[num_contexts:] + reserved_num_blocks = [ + len(block_ids) for block_ids in + kv_cache_manager.get_batch_cache_indices(generation_request_ids) + ] + return list(logical_num_blocks[:num_contexts]) + [ + max(logical, reserved) for logical, reserved in zip( + logical_num_blocks[num_contexts:], reserved_num_blocks) + ] + + def _append_paged_kv_cache( append_key: torch.Tensor, append_value: torch.Tensor, @@ -242,6 +257,9 @@ class FlashInferAttentionMetadata(AttentionMetadata): _ragged_qo_indptr_buf: torch.Tensor = field(init=False) _ragged_kv_indptr_buf: torch.Tensor = field(init=False) _cached_token_lens: torch.Tensor = field(init=False) + # Attention-visible KV lengths. These exclude pages reserved for future speculative tokens when + # the structural page table exposes extra capacity. + _logical_kv_lens: torch.Tensor = field(init=False, repr=False) _plan_params_to_wrappers: Dict[PlanParams, FlashInferWrappers] = field(init=False) @@ -281,6 +299,11 @@ class FlashInferAttentionMetadata(AttentionMetadata): _uses_full_draft_page_table: bool = field(init=False, default=False, repr=False) + # Whether target generation rows expose reserved speculative pages while keeping their + # attention-visible lengths in `_logical_kv_lens`. + _uses_full_generation_page_table: bool = field(init=False, + default=False, + repr=False) def needs_plan(self, plan_params: PlanParams) -> bool: if plan_params not in self._plan_params_to_wrappers: @@ -307,9 +330,7 @@ def get_decode_wrapper( raise ValueError( "FlashInfer draft metadata views require the trtllm-gen " "decode backend.") - num_seqs = self.num_seqs - result._kv_lens_buffer[:num_seqs].copy_( - self._draft_kv_runtime_lens[:num_seqs]) + self._publish_decode_wrapper_kv_lens(result) return result def get_ragged_prefill_wrapper( @@ -509,9 +530,12 @@ def _do_plan_mla_decode(self, plan_params: MLAPlanParams) -> None: qo_indptr = self._qo_indptr[num_ctx:num_ctx + num_gen + 1] - self._qo_indptr[num_ctx] - num_pages_per_seq = kv_indptr[1:] - kv_indptr[:-1] - kv_len_arr = (num_pages_per_seq - - 1) * plan_params.page_size + kv_last_page + if self._uses_full_generation_page_table: + kv_len_arr = self._logical_kv_lens[num_ctx:num_ctx + num_gen] + else: + num_pages_per_seq = kv_indptr[1:] - kv_indptr[:-1] + kv_len_arr = (num_pages_per_seq - + 1) * plan_params.page_size + kv_last_page self._mla_decode_wrapper.plan( qo_indptr, @@ -767,6 +791,69 @@ def update_shared_kv_draft_lengths( num_accepted_tokens[num_contexts:num_seqs]) self._update_draft_kv_lengths() + def apply_spec_decode_kv_lens_offsets( + self, + offsets: torch.Tensor, + num_generations: int, + tokens_per_generation: int, + *, + num_chunked_contexts: int = 0, + restore: bool = False, + ) -> None: + """Apply overlap-scheduler corrections to FlashInfer's live KV state.""" + if self._is_shared_kv_draft_view or self._is_separate_kv_draft_view: + raise RuntimeError( + "Speculative KV offsets must be applied to target metadata") + if not self._uses_full_generation_page_table: + return + if num_chunked_contexts == 0 and num_generations == 0: + return + + direction = -1 if restore else 1 + num_contexts = self.num_contexts + if num_chunked_contexts > 0: + row_slice = slice(num_contexts - num_chunked_contexts, num_contexts) + runtime_offsets = offsets[:num_chunked_contexts] + num_runtime_tokens = num_chunked_contexts * tokens_per_generation + token_slice = slice(self.num_ctx_tokens - num_runtime_tokens, + self.num_ctx_tokens) + else: + row_slice = slice(num_contexts, num_contexts + num_generations) + runtime_offsets = offsets[:num_generations] + num_runtime_tokens = num_generations * tokens_per_generation + token_slice = slice(self.num_ctx_tokens, + self.num_ctx_tokens + num_runtime_tokens) + + self._cached_token_lens[row_slice].add_(runtime_offsets, + alpha=direction) + self._logical_kv_lens[row_slice].add_(runtime_offsets, alpha=direction) + token_offsets = runtime_offsets.repeat_interleave(tokens_per_generation) + self._positions[token_slice].add_(token_offsets, alpha=direction) + + for wrappers in self._plan_params_to_wrappers.values(): + self._publish_decode_wrapper_kv_lens(wrappers.decode_wrapper) + + def _publish_decode_wrapper_kv_lens(self, decode_wrapper) -> None: + """Publish device-logical lengths to a trtllm-gen decode wrapper.""" + kv_lens_buffer = getattr(decode_wrapper, "_kv_lens_buffer", None) + if kv_lens_buffer is None or self.num_generations == 0: + return + + if self._is_shared_kv_draft_view or self._is_separate_kv_draft_view: + kv_lens_buffer[:self.num_generations].copy_( + self._draft_kv_runtime_lens[:self.num_generations]) + return + if not self._uses_full_generation_page_table: + return + + start = self.num_contexts + end = start + self.num_generations + torch.add( + self._cached_token_lens[start:end], + self.seq_lens_kv_cuda[start:end], + out=kv_lens_buffer[:self.num_generations], + ) + def _prepare_full_draft_page_table(self) -> None: """Expose every allocated draft page and use device KV lengths.""" if self._uses_full_draft_page_table: @@ -886,6 +973,13 @@ def _post_init_with_buffers(self, buffers) -> None: self._cached_token_lens = torch.empty((self.max_num_requests, ), dtype=torch.int, device='cuda') + self._logical_kv_lens = self.get_empty( + buffers, + (self.max_num_requests, ), + dtype=torch.int, + cache_name="_logical_kv_lens", + capture_graph=capture_graph, + ) self._draft_kv_runtime_lens = self.get_empty( buffers, (self.max_num_requests, ), @@ -909,6 +1003,7 @@ def _post_init_with_buffers(self, buffers) -> None: self._host_pool_indices: Dict[int, torch.Tensor] = {} self._host_paged_kv_indices: Optional[torch.Tensor] = None self._host_paged_kv_indptr_decode: Optional[torch.Tensor] = None + self._uses_full_generation_page_table = False self._max_num_blocks_per_seq = 0 # VSWA (Variable Sliding Window Attention): models with per-layer @@ -1391,12 +1486,33 @@ def _to_int32_tensor(arr: np.ndarray) -> torch.Tensor: else: self.num_ctx_cached_tokens = 0 - # Number of tokens needed in the KV cache for each sequence after the - # next pass. Kept on the host: every consumer below needs host values, - # so a device-side computation would force a sync per step. + # Number of tokens needed in the KV cache for each sequence after the next pass. + # Kept on the host: every consumer below needs host values, so a device-side computation + # would force a sync per step. These logical block counts describe committed history, not + # all pages that the KV manager has reserved for fixed-width speculative appends. kv_lens_host = np.asarray(num_cached_tokens_per_seq, dtype=np.int64) + self.seq_lens_kv.numpy() - num_blocks = (kv_lens_host + self.page_size - 1) // self.page_size + logical_num_blocks = ((kv_lens_host + self.page_size - 1) // + self.page_size) + num_blocks = logical_num_blocks + use_full_generation_page_table = ( + self.kv_cache_params.use_full_generation_page_table) + self._uses_full_generation_page_table = use_full_generation_page_table + if use_full_generation_page_table: + # Generation rows need every reserved page to keep a speculative append addressable + # across a page boundary. Treating that wider table as the logical sequence would make + # FlashInfer attend to uncommitted draft slots, so plans and positions use the separate + # device logical lengths populated below. + assert self.request_ids is not None + num_blocks = np.asarray( + _get_page_table_num_blocks( + self.kv_cache_manager, + self.request_ids, + logical_num_blocks.tolist(), + self.num_contexts, + ), + dtype=np.int64, + ) self.num_blocks = num_blocks.tolist() assert self.request_ids is not None @@ -1464,9 +1580,9 @@ def _to_int32_tensor(arr: np.ndarray) -> torch.Tensor: self._vswa_active_pool_id = primary_pool_id # number of tokens in the last cache block used by each sequence, - # derived on the host so no GPU arithmetic or sync is needed. + # derived from the logical rather than reservation-width page count. paged_kv_last_page_len = _to_int32_tensor(kv_lens_host - - (num_blocks - 1) * + (logical_num_blocks - 1) * self.page_size) self._paged_kv_last_page_len[:paged_kv_last_page_len.size(0)].copy_( paged_kv_last_page_len, non_blocking=True) @@ -1509,11 +1625,23 @@ def _to_int32_tensor(arr: np.ndarray) -> torch.Tensor: # For cross attention, num_tokens is 0 during decode, and we don't need to update kv cache. if self.num_tokens > 0: + if use_full_generation_page_table: + # The page-table indptr now describes addressable capacity. Deriving positions from + # it would count reserved pages as committed KV and potentially shift appends beyond + # the true sequence end. + logical_kv_lens = _to_int32_tensor(kv_lens_host) + self._logical_kv_lens[:logical_kv_lens.numel()].copy_( + logical_kv_lens, non_blocking=True) + position_kv_lens = self._logical_kv_lens[:self.num_seqs] + else: + position_kv_lens = flashinfer.get_seq_lens( + self.paged_kv_indptr, + self.paged_kv_last_page_len, + self.page_size, + ) batch_indices, positions = flashinfer.get_batch_indices_positions( self.kv_indptr, - flashinfer.get_seq_lens(self.paged_kv_indptr, - self.paged_kv_last_page_len, - self.page_size), + position_kv_lens, self.num_tokens, ) self._batch_indices[:batch_indices.size(0)].copy_(batch_indices, @@ -1858,6 +1986,11 @@ def prefill_plan(): def decode_plan(): assert decode_wrapper is not None + if (self._uses_full_generation_page_table + and decode_wrapper._backend != "trtllm-gen"): + raise ValueError( + "Reservation-width FlashInfer page tables require the " + "trtllm-gen decode backend's independent KV lengths.") # Host int32 indptr (retained by prepare, which always runs # before plans): flashinfer moves it to the device itself, and # its indptr.cpu()/get_seq_lens calls stay free of D2H syncs. @@ -1885,6 +2018,7 @@ def decode_plan(): o_data_type=o_dtype, block_tables=block_tables, ) + self._publish_decode_wrapper_kv_lens(decode_wrapper) # Must sync after append_paged_kv_cache and before plan. torch.cuda.current_stream().synchronize() diff --git a/tensorrt_llm/_torch/metadata.py b/tensorrt_llm/_torch/metadata.py index fd076cc4cab0..40962845464a 100644 --- a/tensorrt_llm/_torch/metadata.py +++ b/tensorrt_llm/_torch/metadata.py @@ -1,3 +1,6 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + from dataclasses import dataclass from enum import Enum from typing import List, Optional @@ -30,6 +33,10 @@ class KVCacheParams: # The number of extra kv for draft tokens num_extra_kv_tokens: Optional[int] = 0 + # Whether generation page tables expose all reserved pages. Backends that + # use this must track the logical KV length separately. + use_full_generation_page_table: bool = False + class CacheType(Enum): # Linear KV cache stores all the cached tokens of a sequence in a single page. diff --git a/tensorrt_llm/_torch/models/modeling_gemma4.py b/tensorrt_llm/_torch/models/modeling_gemma4.py index 1834217b57bc..9acd5a5e2b3f 100644 --- a/tensorrt_llm/_torch/models/modeling_gemma4.py +++ b/tensorrt_llm/_torch/models/modeling_gemma4.py @@ -14,6 +14,7 @@ # limitations under the License. """TensorRT-LLM PyTorch backend implementation for Gemma4 text model.""" +import copy import dataclasses import math from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple, Union @@ -59,7 +60,11 @@ from ..modules.rms_norm import RMSNorm from ..speculative.interface import SpecMetadata from ..utils import ActivationType, Fp4QuantizedTensor, is_torch_compiling -from .modeling_speculative import SpecDecOneEngineForCausalLM, _slice_spec_position_ids +from .modeling_speculative import ( + _SPECULATIVE_POSITION_HEADROOM, + SpecDecOneEngineForCausalLM, + _slice_spec_position_ids, +) from .modeling_utils import DecoderModel, DecoderModelForCausalLM, register_auto_model if TYPE_CHECKING: @@ -78,6 +83,17 @@ from transformers import Gemma4TextConfig # noqa: E402 +def _gemma4_rope_max_positions(model_config: ModelConfig) -> int: + """Size RoPE for transient speculative positions beyond max_seq_len.""" + max_positions = model_config.pretrained_config.max_position_embeddings + spec_config = model_config.spec_config + if spec_config is not None: + # Overlap can have one pending verification width, while the current target and shared, + # Q-only assistant consume the next one. + max_positions += 2 * spec_config.tokens_per_gen_step + return max_positions + + # --------------------------------------------------------------------------- # Scaled embedding (reused from Gemma3 pattern) # --------------------------------------------------------------------------- @@ -226,7 +242,7 @@ def __init__( # Build RoPE params per layer type rope_params = RopeParams() - rope_params.max_positions = config.max_position_embeddings + rope_params.max_positions = _gemma4_rope_max_positions(model_config) if is_sliding: # Sliding: default RoPE, theta=10K, full rotation rope_config = ( @@ -1597,9 +1613,16 @@ class Gemma4AssistantForCausalLM(DecoderModelForCausalLM[Gemma4TextModel, Gemma4 def __init__(self, model_config: ModelConfig): assistant_config = model_config.pretrained_config + assistant_text_config = copy.deepcopy(assistant_config.text_config) + # The assistant is deliberately built without `spec_config` to avoid recursive speculative + # initialization. Carry only its fixed physical position headroom through the internal model + # attributes instead. + assistant_text_config.max_position_embeddings += model_config.extra_attrs.get( + _SPECULATIVE_POSITION_HEADROOM, 0 + ) text_model_config = dataclasses.replace( model_config, - pretrained_config=assistant_config.text_config, + pretrained_config=assistant_text_config, spec_config=None, ) super().__init__( diff --git a/tensorrt_llm/_torch/models/modeling_speculative.py b/tensorrt_llm/_torch/models/modeling_speculative.py index 0715c39580a3..a270b59b7f2f 100755 --- a/tensorrt_llm/_torch/models/modeling_speculative.py +++ b/tensorrt_llm/_torch/models/modeling_speculative.py @@ -45,6 +45,8 @@ from .modeling_utils import (DecoderModel, DecoderModelForCausalLM, TModel, get_model_architecture, register_auto_model) +_SPECULATIVE_POSITION_HEADROOM = "_speculative_position_headroom" + def _ensure_draft_vocab_size(config: PretrainedConfig) -> None: if hasattr(config, @@ -2552,7 +2554,11 @@ def __init__(self, moe_max_num_tokens=model_config.moe_max_num_tokens) self.draft_config.quant_config.kv_cache_quant_algo = \ model_config.quant_config.kv_cache_quant_algo - self.draft_config.extra_attrs = model_config.extra_attrs + self.draft_config.extra_attrs = dict( + model_config.extra_attrs) + self.draft_config.extra_attrs[ + _SPECULATIVE_POSITION_HEADROOM] = ( + 2 * spec_config.tokens_per_gen_step) elif spec_config.spec_dec_mode.is_external_drafter(): self.draft_config = ModelConfig.from_pretrained( diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 3da7a73a7e9e..e5c096bae84d 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -3755,6 +3755,16 @@ def get_max_num_sequences(self) -> int: num_batches = self.mapping.pp_size return num_batches * self.batch_size + def _should_use_full_generation_page_table( + self, spec_config: Optional[DecodingBaseConfig], + attn_metadata: AttentionMetadata) -> bool: + """Return whether overlap decode needs every reserved generation page.""" + # FlashInfer metadata owns the optional device-side KV-length correction used with this + # wider page table. + return (self.enable_spec_decode and not self._disable_overlap_scheduler + and getattr(spec_config, '_use_shared_kv_cache', False) + and hasattr(attn_metadata, 'apply_spec_decode_kv_lens_offsets')) + def _preprocess_inputs(self, inputs: Dict[str, Any]): """ Make some changes to the device inputs and avoid blocking the async data transfer @@ -3806,6 +3816,17 @@ def _preprocess_inputs(self, inputs: Dict[str, Any]): previous_kv_lens_offsets_cuda[:num_gen_requests] ) inputs['attn_metadata'].on_update_kv_lens() + # TRTLLM uses `kv_lens_cuda` above; FlashInfer exposes this backend-specific + # correction without coupling the engine to its metadata type. + elif hasattr(inputs['attn_metadata'], + 'apply_spec_decode_kv_lens_offsets'): + inputs['attn_metadata'].apply_spec_decode_kv_lens_offsets( + self.previous_kv_lens_offsets_cuda, + num_gen_requests, + self.get_runtime_tokens_per_gen_step( + self.runtime_draft_len), + num_chunked_contexts=num_chunked_ctx_requests, + ) if self.guided_decoder is not None: self.guided_decoder.token_event.record() @@ -3853,6 +3874,18 @@ def _postprocess_inputs(self, inputs: Dict[str, Any]): self. previous_kv_lens_offsets_cuda[:num_gen_requests] ) + # Restore the FlashInfer-specific logical KV lengths through the same optional hook + # used by `_preprocess_inputs`. + elif hasattr(inputs['attn_metadata'], + 'apply_spec_decode_kv_lens_offsets'): + inputs['attn_metadata'].apply_spec_decode_kv_lens_offsets( + self.previous_kv_lens_offsets_cuda, + num_gen_requests, + self.get_runtime_tokens_per_gen_step( + self.runtime_draft_len), + num_chunked_contexts=num_chunked_ctx_requests, + restore=True, + ) def _get_all_rank_num_tokens(self, attn_metadata: AttentionMetadata): if self.enable_attention_dp: @@ -4556,7 +4589,10 @@ def _prepare_incremental_update_metadata( attn_metadata.kv_cache_params = KVCacheParams( use_cache=True, num_cached_tokens_per_seq=num_cached_tokens_per_seq, - num_extra_kv_tokens=get_num_extra_kv_tokens(spec_config)) + num_extra_kv_tokens=get_num_extra_kv_tokens(spec_config), + use_full_generation_page_table=( + self._should_use_full_generation_page_table( + spec_config, attn_metadata))) attn_metadata.kv_cache_manager = kv_cache_manager attn_metadata.prepare() @@ -6217,7 +6253,10 @@ def previous_seq_slots_device(): attn_metadata.kv_cache_params = KVCacheParams( use_cache=True, num_cached_tokens_per_seq=num_cached_tokens_per_seq, - num_extra_kv_tokens=get_num_extra_kv_tokens(spec_config)) + num_extra_kv_tokens=get_num_extra_kv_tokens(spec_config), + use_full_generation_page_table=( + self._should_use_full_generation_page_table( + spec_config, attn_metadata))) attn_metadata.kv_cache_manager = kv_cache_manager if hasattr(self.model.model_config.pretrained_config, 'chunk_size'): diff --git a/tests/integration/defs/accuracy/test_llm_api_pytorch_multimodal.py b/tests/integration/defs/accuracy/test_llm_api_pytorch_multimodal.py index 192bdc72bc48..6ac701374cbd 100644 --- a/tests/integration/defs/accuracy/test_llm_api_pytorch_multimodal.py +++ b/tests/integration/defs/accuracy/test_llm_api_pytorch_multimodal.py @@ -324,9 +324,6 @@ def test_nvfp4(self): max_batch_size=16, kv_cache_config=self.kv_cache_config, enable_chunked_prefill=True, - # Shared-KV MTP overlap can expose too few FlashInfer pages and cause an illegal access - # in `AppendPagedKVCache`. Re-enable this after the overlap-MTP KV accounting fix lands. - disable_overlap_scheduler=True, speculative_config=MTPDecodingConfig( max_draft_len=3, mtp_eagle_one_model=True, diff --git a/tests/unittest/_torch/attention/test_flashinfer_attention.py b/tests/unittest/_torch/attention/test_flashinfer_attention.py index fe00e3d69d07..0c037928d248 100644 --- a/tests/unittest/_torch/attention/test_flashinfer_attention.py +++ b/tests/unittest/_torch/attention/test_flashinfer_attention.py @@ -1,7 +1,11 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + import random import unittest from collections import defaultdict from dataclasses import dataclass +from types import SimpleNamespace from typing import List, Optional, Union from unittest import mock @@ -13,7 +17,8 @@ FlashInferAttentionMetadata) from tensorrt_llm._torch.attention_backend import \ flashinfer as flashinfer_backend -from tensorrt_llm._torch.attention_backend.flashinfer import PlanParams +from tensorrt_llm._torch.attention_backend.flashinfer import ( + FlashInferWrappers, PlanParams) from tensorrt_llm._torch.attention_backend.interface import \ PredefinedAttentionMask from tensorrt_llm._torch.metadata import KVCacheParams @@ -66,6 +71,168 @@ class CUDAGraphTestScenario: class TestFlashInferAttention(unittest.TestCase): + def test_generation_page_table_uses_reserved_block_count(self): + manager = SimpleNamespace(get_batch_cache_indices=mock.Mock( + return_value=[list(range(325))])) + + self.assertEqual( + flashinfer_backend._get_page_table_num_blocks(manager, [98, 99], + [3, 324], + num_contexts=1), + [3, 325], + ) + manager.get_batch_cache_indices.assert_called_once_with([99]) + + def test_generation_page_table_keeps_logical_positions(self): + if not torch.cuda.is_available(): + self.skipTest("CUDA is required for FlashInfer metadata") + + kv_cache_manager = KVCacheManager( + KvCacheConfig(max_tokens=256), + tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF, + num_layers=1, + num_kv_heads=1, + head_dim=128, + tokens_per_block=32, + max_seq_len=64, + max_batch_size=1, + mapping=Mapping(world_size=1, tp_size=1, rank=0), + dtype=tensorrt_llm.bindings.DataType.BF16, + ) + try: + kv_cache_manager.add_dummy_requests([0], [32], + is_gen=True, + max_num_draft_tokens=3) + reserved_blocks = kv_cache_manager.get_batch_cache_indices([0]) + self.assertEqual(len(reserved_blocks[0]), 2) + + metadata = FlashInferAttentionMetadata( + seq_lens=torch.full((1, ), 4, dtype=torch.int32), + num_contexts=0, + kv_cache_params=KVCacheParams( + use_cache=True, + num_cached_tokens_per_seq=[28], + use_full_generation_page_table=True, + ), + max_num_requests=1, + max_num_tokens=4, + kv_cache_manager=kv_cache_manager, + request_ids=[0], + ) + metadata.prepare() + + self.assertEqual(metadata.num_blocks, [2]) + torch.testing.assert_close( + metadata.paged_kv_indptr_decode[:2], + torch.tensor([0, 2], dtype=torch.int32, device="cuda"), + ) + torch.testing.assert_close( + metadata._paged_kv_last_page_len[:1], + torch.tensor([32], dtype=torch.int32, device="cuda"), + ) + torch.testing.assert_close( + metadata._logical_kv_lens[:1], + torch.tensor([32], dtype=torch.int32, device="cuda"), + ) + torch.testing.assert_close( + metadata.positions, + torch.tensor([28, 29, 30, 31], dtype=torch.int32, + device="cuda"), + ) + finally: + kv_cache_manager.shutdown() + + def test_spec_decode_offsets_update_append_and_decode_lengths(self): + if not torch.cuda.is_available(): + self.skipTest("CUDA is required for FlashInfer metadata") + + kv_cache_manager = KVCacheManager( + KvCacheConfig(max_tokens=256), + tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF, + num_layers=1, + num_kv_heads=1, + head_dim=128, + tokens_per_block=32, + max_seq_len=64, + max_batch_size=2, + mapping=Mapping(world_size=1, tp_size=1, rank=0), + dtype=tensorrt_llm.bindings.DataType.BF16, + ) + try: + metadata = FlashInferAttentionMetadata( + seq_lens=torch.full((2, ), 4, dtype=torch.int32), + num_contexts=0, + kv_cache_params=KVCacheParams(use_cache=True), + max_num_requests=2, + max_num_tokens=8, + kv_cache_manager=kv_cache_manager, + ) + cached_token_lens = torch.tensor([31, 62], + dtype=torch.int32, + device="cuda") + positions = torch.tensor([31, 32, 33, 34, 62, 63, 64, 65], + dtype=torch.int32, + device="cuda") + metadata._cached_token_lens[:2].copy_(cached_token_lens) + metadata._logical_kv_lens[:2].copy_(cached_token_lens + 4) + metadata._uses_full_generation_page_table = True + metadata._positions[:8].copy_(positions) + kv_lens_buffer = torch.tensor([35, 66], + dtype=torch.int32, + device="cuda") + metadata._plan_params_to_wrappers = { + object(): + FlashInferWrappers( + is_planned=True, + decode_wrapper=SimpleNamespace( + _kv_lens_buffer=kv_lens_buffer), + ) + } + offsets = torch.tensor([-3, -1], dtype=torch.int32, device="cuda") + + metadata.apply_spec_decode_kv_lens_offsets( + offsets, + num_generations=2, + tokens_per_generation=4, + ) + + torch.testing.assert_close( + metadata._cached_token_lens[:2], + torch.tensor([28, 61], dtype=torch.int32, device="cuda"), + ) + torch.testing.assert_close( + metadata._logical_kv_lens[:2], + torch.tensor([32, 65], dtype=torch.int32, device="cuda"), + ) + torch.testing.assert_close( + metadata._positions[:8], + torch.tensor([28, 29, 30, 31, 61, 62, 63, 64], + dtype=torch.int32, + device="cuda"), + ) + torch.testing.assert_close( + kv_lens_buffer, + torch.tensor([32, 65], dtype=torch.int32, device="cuda"), + ) + + metadata.apply_spec_decode_kv_lens_offsets( + offsets, + num_generations=2, + tokens_per_generation=4, + restore=True, + ) + torch.testing.assert_close(metadata._cached_token_lens[:2], + cached_token_lens) + torch.testing.assert_close(metadata._logical_kv_lens[:2], + cached_token_lens + 4) + torch.testing.assert_close(metadata._positions[:8], positions) + torch.testing.assert_close( + kv_lens_buffer, + torch.tensor([35, 66], dtype=torch.int32, device="cuda"), + ) + finally: + kv_cache_manager.shutdown() + def test_separate_kv_draft_metadata_uses_draft_manager(self): if not torch.cuda.is_available(): self.skipTest("CUDA is required for FlashInfer metadata") diff --git a/tests/unittest/_torch/modeling/test_modeling_gemma4.py b/tests/unittest/_torch/modeling/test_modeling_gemma4.py index 5eff2c644555..9745126f0ee2 100644 --- a/tests/unittest/_torch/modeling/test_modeling_gemma4.py +++ b/tests/unittest/_torch/modeling/test_modeling_gemma4.py @@ -46,6 +46,7 @@ Gemma4TextScaledWordEmbedding, ) from tensorrt_llm._utils import is_sm_100f +from tensorrt_llm.llmapi.llm_args import MTPDecodingConfig from tensorrt_llm.mapping import Mapping if TYPE_CHECKING: @@ -381,6 +382,26 @@ def test_full_rope_params(self): expected_dim = int(config.global_head_dim * 0.25) self.assertEqual(rope.dim, expected_dim) + def test_rope_params_include_speculative_headroom(self): + """RoPE must cover draft positions beyond the logical sequence limit.""" + model_config = _make_model_config(GEMMA4_SMALL_CONFIG) + spec_config = MTPDecodingConfig(max_draft_len=3) + spec_config._use_shared_kv_cache = True + model_config.spec_config = spec_config + model_config.attn_backend = "FLASHINFER" + + expected_max_positions = GEMMA4_SMALL_CONFIG["max_position_embeddings"] + 2 * 4 + for layer_idx, is_sliding in ((0, True), (5, False)): + with self.subTest(is_sliding=is_sliding): + attn = Gemma4Attention( + model_config, + layer_idx=layer_idx, + is_sliding=is_sliding, + ) + self.assertEqual(attn.pos_embd_params.rope.max_positions, expected_max_positions) + self.assertEqual(attn.rotary_emb.max_positions, expected_max_positions) + self.assertEqual(attn.rotary_emb.rotary_cos_sin.shape[0], expected_max_positions) + def test_num_kv_heads_per_layer_type(self): """Sliding layers use num_key_value_heads, full use num_global_key_value_heads.""" model_config = _make_model_config(GEMMA4_SMALL_CONFIG) @@ -618,9 +639,24 @@ def test_ordered_embedding_combines_vocab_parallel_shards(self): torch.testing.assert_close(actual, expected) def test_assistant_uses_target_kv_sources(self): - assistant = Gemma4AssistantForCausalLM(_make_assistant_model_config()) + model_config = _make_assistant_model_config() + model_config.extra_attrs["_speculative_position_headroom"] = 2 * 4 + assistant = Gemma4AssistantForCausalLM(model_config) self.assertEqual(len(assistant.model.layers), 4) self.assertTrue(all(layer.is_kv_shared_layer for layer in assistant.model.layers)) + self.assertEqual( + assistant.model.model_config.pretrained_config.max_position_embeddings, + GEMMA4_SMALL_CONFIG["max_position_embeddings"] + 2 * 4, + ) + self.assertEqual( + model_config.pretrained_config.text_config.max_position_embeddings, + GEMMA4_SMALL_CONFIG["max_position_embeddings"], + ) + for layer in assistant.model.layers: + self.assertEqual( + layer.self_attn.pos_embd_params.rope.max_positions, + GEMMA4_SMALL_CONFIG["max_position_embeddings"] + 2 * 4, + ) target_config = { **GEMMA4_SMALL_CONFIG, From 9ef9c8109b17d0f922983074f2170d313af36de6 Mon Sep 17 00:00:00 2001 From: William Zhang <133824995+2ez4bz@users.noreply.github.com> Date: Wed, 19 Aug 2026 11:02:09 -0700 Subject: [PATCH 2/2] [None][fix] Replan FlashInfer decode for launch shape * Why? Cached decode plans ignored query length and generation batch size, so speculative decoding could reuse a plan built for an incompatible launch shape. * What? Include both dimensions in the plan parameters, pass the query length to FlashInfer, and reject batches with nonuniform generation query lengths. Signed-off-by: William Zhang <133824995+2ez4bz@users.noreply.github.com> --- .../_torch/attention_backend/flashinfer.py | 29 +++++++- .../attention/test_flashinfer_attention.py | 70 +++++++++++++++++++ 2 files changed, 97 insertions(+), 2 deletions(-) diff --git a/tensorrt_llm/_torch/attention_backend/flashinfer.py b/tensorrt_llm/_torch/attention_backend/flashinfer.py index 129c18c607c2..a1fb669369a0 100644 --- a/tensorrt_llm/_torch/attention_backend/flashinfer.py +++ b/tensorrt_llm/_torch/attention_backend/flashinfer.py @@ -166,8 +166,10 @@ class FlashInferMultiItemParams: @dataclass(kw_only=True, frozen=True) class PlanParams: - """ - Parameters that affect the flashinfer execution plan + """Parameters that affect FlashInfer wrapper planning. + + Include values that change wrapper-owned CUDA graph state, even when a backend derives them + again at runtime. """ num_heads: int @@ -182,6 +184,11 @@ class PlanParams: sm_scale: Optional[float] = None window_left: Optional[int] = None kv_pool_id: Optional[int] = None + # Decode wrappers own persistent graph-visible buffers and counters. The speculative query width + # and generation batch size must distinguish cache entries; reusing a wrapper across either + # dimension can expose stale launch state. + q_len_per_req: int = 1 + num_generations: int = 0 # NB: Some features (multi-item scoring) are only supported with the paged KV-cache wrapper. @@ -1815,6 +1822,20 @@ def plan(self, if q_scaling is not None: sm_scale = 1 / (math.sqrt(head_dim) * q_scaling) + # FlashInfer decode accepts one q_len_per_req. Current paths keep generation widths + # uniform (including padded speculative drafts); guard against malformed metadata or + # future per-request draft widths rather than launching the wrong shape. + q_len_per_req = 1 + if self.num_generations > 0: + generation_seq_lens = self.seq_lens[self. + num_contexts:self.num_contexts + + self.num_generations] + q_len_per_req = int(generation_seq_lens[0]) + if not torch.all(generation_seq_lens == q_len_per_req): + raise ValueError( + "FlashInfer decode requires a uniform query length per " + f"request, but got {generation_seq_lens.tolist()}") + plan_params = PlanParams( num_heads=num_heads, num_kv_heads=num_kv_heads, @@ -1828,6 +1849,8 @@ def plan(self, attention_mask_data=attention_mask_data, multi_item_params=self._multi_item_params, kv_pool_id=getattr(self, "_vswa_active_pool_id", None), + q_len_per_req=q_len_per_req, + num_generations=self.num_generations, ) return self._plan_with_params(plan_params, flashinfer_backend) @@ -2017,6 +2040,8 @@ def decode_plan(): kv_data_type=plan_params.kv_dtype, o_data_type=o_dtype, block_tables=block_tables, + # Keep FlashInfer's recorded graph shape aligned with the wrapper cache key. + q_len_per_req=plan_params.q_len_per_req, ) self._publish_decode_wrapper_kv_lens(decode_wrapper) diff --git a/tests/unittest/_torch/attention/test_flashinfer_attention.py b/tests/unittest/_torch/attention/test_flashinfer_attention.py index 0c037928d248..0b933623acd5 100644 --- a/tests/unittest/_torch/attention/test_flashinfer_attention.py +++ b/tests/unittest/_torch/attention/test_flashinfer_attention.py @@ -83,6 +83,76 @@ def test_generation_page_table_uses_reserved_block_count(self): ) manager.get_batch_cache_indices.assert_called_once_with([99]) + def test_decode_launch_shape_is_part_of_plan_params(self): + if not torch.cuda.is_available(): + self.skipTest("CUDA is required for FlashInfer metadata") + + metadata = FlashInferAttentionMetadata( + seq_lens=torch.tensor([1, 1], dtype=torch.int32), + num_contexts=0, + kv_cache_manager=None, + request_ids=[0, 1], + max_num_requests=3, + max_num_tokens=18, + ) + + def return_plan_params(plan_params, _flashinfer_backend): + return plan_params + + with mock.patch.object( + metadata, + "_plan_with_params", + side_effect=return_plan_params, + ): + single_token_plan = metadata.plan( + num_heads=32, + num_kv_heads=4, + head_dim=512, + q_dtype=torch.float8_e4m3fn, + kv_dtype=torch.float8_e4m3fn, + attention_mask_type=AttentionMaskType.causal.value, + flashinfer_backend="trtllm-gen", + ) + metadata.seq_lens = torch.tensor([6, 6], dtype=torch.int32) + multi_token_plan = metadata.plan( + num_heads=32, + num_kv_heads=4, + head_dim=512, + q_dtype=torch.float8_e4m3fn, + kv_dtype=torch.float8_e4m3fn, + attention_mask_type=AttentionMaskType.causal.value, + flashinfer_backend="trtllm-gen", + ) + metadata.seq_lens = torch.tensor([6, 6, 6], dtype=torch.int32) + larger_batch_plan = metadata.plan( + num_heads=32, + num_kv_heads=4, + head_dim=512, + q_dtype=torch.float8_e4m3fn, + kv_dtype=torch.float8_e4m3fn, + attention_mask_type=AttentionMaskType.causal.value, + flashinfer_backend="trtllm-gen", + ) + + self.assertEqual(single_token_plan.q_len_per_req, 1) + self.assertEqual(multi_token_plan.q_len_per_req, 6) + self.assertNotEqual(single_token_plan, multi_token_plan) + self.assertEqual(multi_token_plan.num_generations, 2) + self.assertEqual(larger_batch_plan.num_generations, 3) + self.assertNotEqual(multi_token_plan, larger_batch_plan) + + metadata.seq_lens = torch.tensor([6, 5], dtype=torch.int32) + with self.assertRaisesRegex(ValueError, "uniform query length"): + metadata.plan( + num_heads=32, + num_kv_heads=4, + head_dim=512, + q_dtype=torch.float8_e4m3fn, + kv_dtype=torch.float8_e4m3fn, + attention_mask_type=AttentionMaskType.causal.value, + flashinfer_backend="trtllm-gen", + ) + def test_generation_page_table_keeps_logical_positions(self): if not torch.cuda.is_available(): self.skipTest("CUDA is required for FlashInfer metadata")