diff --git a/tensorrt_llm/_torch/attention_backend/flashinfer.py b/tensorrt_llm/_torch/attention_backend/flashinfer.py index d49ab084a128..a1fb669369a0 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, @@ -151,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 @@ -167,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. @@ -242,6 +264,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 +306,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 +337,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 +537,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 +798,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 +980,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 +1010,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 +1493,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 +1587,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 +1632,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, @@ -1687,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, @@ -1700,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) @@ -1858,6 +2009,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. @@ -1884,7 +2040,10 @@ 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) # 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..0b933623acd5 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,238 @@ 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_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") + + 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,