From ab28024d1874b1c7fe5fe439362f5c4a4beee086 Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Wed, 19 Aug 2026 22:46:16 -0700 Subject: [PATCH 1/9] [None][fix] Handle evicted SWA pages in FlashInfer attention Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- .../_torch/attention_backend/flashinfer.py | 28 +++++- .../_torch/modeling/test_modeling_gemma4.py | 97 ++++++++++++++++++- 2 files changed, 120 insertions(+), 5 deletions(-) diff --git a/tensorrt_llm/_torch/attention_backend/flashinfer.py b/tensorrt_llm/_torch/attention_backend/flashinfer.py index a1fb669369a0..c949137e2be8 100644 --- a/tensorrt_llm/_torch/attention_backend/flashinfer.py +++ b/tensorrt_llm/_torch/attention_backend/flashinfer.py @@ -37,6 +37,7 @@ from tensorrt_llm.functional import AttentionMaskType from tensorrt_llm.logger import logger from tensorrt_llm.models.modeling_utils import QuantConfig +from tensorrt_llm.runtime.kv_cache_manager_v2._common import BAD_PAGE_INDEX from ..metadata import KVCacheParams from ..utils import get_global_attrs, get_model_extra_attrs, torch_multi_arange @@ -579,6 +580,20 @@ def get_paged_kv_indices_for_layer(self, layer_idx: int) -> torch.Tensor: total_blocks = self.num_generation_blocks + self.num_context_blocks return self._vswa_pool_indices_cache[pool_id][:total_blocks] + def _sanitize_swa_page_indices(self, page_indices: torch.Tensor, + layer_idx: int) -> None: + """Replace evicted SWA pages with a safe in-range page index.""" + window_vec = getattr(self.kv_cache_manager, 'max_attention_window_vec', + None) + if not window_vec or window_vec[layer_idx % len(window_vec)] is None: + return + + # KVCacheManagerV2 marks evicted out-of-window pages with -1. + # FlashInfer may dereference page IDs before applying window_left, so + # keep masked positions in range. The SWA mask excludes their values + # from the attention result. + page_indices.masked_fill_(page_indices == BAD_PAGE_INDEX, 0) + def swap_paged_kv_indices_for_layer(self, layer_idx: int) -> None: """Copy pool-specific page indices into the shared buffer. @@ -1542,8 +1557,18 @@ def _to_int32_tensor(arr: np.ndarray) -> torch.Tensor: self.num_generation_blocks = sum(self.num_blocks[self.num_contexts:]) # indices of used cache blocks for each sequence + primary_layer_idx = None + if self._vswa_layer_to_pool is not None: + primary_pool_id = self._vswa_layer_to_pool.get(0, 0) + primary_layer_idx = self._vswa_pool_to_rep_layer[primary_pool_id] + else: + layer_offsets = getattr(self.kv_cache_manager, 'layer_offsets', {}) + primary_layer_idx = next(iter(layer_offsets), None) + paged_kv_indices = self.kv_cache_manager.get_batch_cache_indices_flat( - self.request_ids, self.num_blocks) + self.request_ids, self.num_blocks, layer_idx=primary_layer_idx) + if primary_layer_idx is not None: + self._sanitize_swa_page_indices(paged_kv_indices, primary_layer_idx) self._paged_kv_indices[:paged_kv_indices.size(0)].copy_( paged_kv_indices, non_blocking=True) @@ -1579,6 +1604,7 @@ def _to_int32_tensor(arr: np.ndarray) -> torch.Tensor: pool_indices = \ self.kv_cache_manager.get_batch_cache_indices_flat( self.request_ids, self.num_blocks, layer_idx=rep_layer) + self._sanitize_swa_page_indices(pool_indices, rep_layer) buf = getattr(self, f'_vswa_pool_buf_{pool_id}') buf[:pool_indices.size(0)].copy_(pool_indices, non_blocking=True) diff --git a/tests/unittest/_torch/modeling/test_modeling_gemma4.py b/tests/unittest/_torch/modeling/test_modeling_gemma4.py index 9745126f0ee2..26a57427de0c 100644 --- a/tests/unittest/_torch/modeling/test_modeling_gemma4.py +++ b/tests/unittest/_torch/modeling/test_modeling_gemma4.py @@ -48,6 +48,7 @@ from tensorrt_llm._utils import is_sm_100f from tensorrt_llm.llmapi.llm_args import MTPDecodingConfig from tensorrt_llm.mapping import Mapping +from tensorrt_llm.runtime.kv_cache_manager_v2._common import BAD_PAGE_INDEX if TYPE_CHECKING: from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2 @@ -962,17 +963,16 @@ def _build_gemma4_kv_cache_manager( # Set per-layer max_attention_window when head_dim or kv_heads differ # across layers, so V2 creates separate pool groups for different page - # sizes. ``max_seq_len - 1`` on sliding layers prevents V2 block - # eviction that would cause FlashInfer page index OOB when kv_lens - # exceeds sliding_window. + # sizes. sliding_window = getattr(config, "sliding_window", None) max_attn_window = None needs_vswa = isinstance(head_dim, list) and len(set(head_dim)) > 1 if not needs_vswa: needs_vswa = isinstance(num_kv_heads, list) and len(set(num_kv_heads)) > 1 if needs_vswa and sliding_window: + swa_window = min(sliding_window, max_seq_len - 1) max_attn_window = [ - max_seq_len - 1 if lt == "sliding_attention" else max_seq_len for lt in layer_types + swa_window if lt == "sliding_attention" else max_seq_len for lt in layer_types ] kv_cache_config = KvCacheConfigV2( @@ -3559,6 +3559,95 @@ def test_cuda_graph_decode_26b_like(self): """26B-like: GQA=2, K=V, hd=256/512.""" self._run_cuda_graph_real_headdim(deepcopy(GEMMA4_26B_REAL_DIMS_CONFIG), "26B") + @torch.no_grad() + @unittest.mock.patch( + "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None + ) + def test_cuda_graph_decode_with_evicted_swa_pages(self): + """FA2 CUDA graph decode safely ignores evicted SWA pages.""" + config_dict = deepcopy(GEMMA4_26B_REAL_DIMS_CONFIG) + config_dict["sliding_window"] = 1024 + config_dict["max_position_embeddings"] = 4096 + config = Gemma4TextConfig(**config_dict) + + tokens_per_block = 32 + cached_tokens = 2046 + kv_cache_manager = self._get_kv_cache_manager( + config, + num_blocks=80, + tokens_per_block=tokens_per_block, + batch_size=1, + ) + self.addCleanup(kv_cache_manager.shutdown) + + requests = kv_cache_manager.add_dummy_requests([0], [cached_tokens + 1], is_gen=True) + self.assertIsNotNone(requests) + torch.nn.init.normal_(kv_cache_manager.get_buffers(0)) + + num_blocks = (cached_tokens + 1 + tokens_per_block - 1) // tokens_per_block + raw_indices = kv_cache_manager.get_batch_cache_indices_flat([0], [num_blocks], layer_idx=0) + self.assertIn(BAD_PAGE_INDEX, raw_indices.tolist()) + + metadata = FlashInferAttentionMetadata( + seq_lens=torch.ones(1, dtype=torch.int), + num_contexts=0, + is_cuda_graph=True, + kv_cache_params=KVCacheParams( + use_cache=True, + num_cached_tokens_per_seq=[cached_tokens], + ), + workspace_buffer=torch.empty( + _FLASHINFER_WORKSPACE_BYTES, dtype=torch.uint8, device="cuda" + ), + max_num_requests=1, + max_num_tokens=4096, + kv_cache_manager=kv_cache_manager, + request_ids=[0], + ) + metadata.prepare() + sanitized_indices = metadata.get_paged_kv_indices_for_layer(0) + self.assertNotIn(BAD_PAGE_INDEX, sanitized_indices.cpu().tolist()) + + layer = FlashInferAttention( + layer_idx=0, + num_heads=config.num_attention_heads, + head_dim=config.head_dim, + num_kv_heads=config.num_key_value_heads, + q_scaling=1.0 / math.sqrt(config.head_dim), + flashinfer_backend="fa2", + ) + query = torch.randn( + 1, + config.num_attention_heads * config.head_dim, + dtype=config.torch_dtype, + device="cuda", + ) + + eager_output = None + for _ in range(2): + eager_output = layer.forward( + query, + None, + None, + metadata, + attention_window_size=config.sliding_window, + ) + + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + graph_output = layer.forward( + query, + None, + None, + metadata, + attention_window_size=config.sliding_window, + ) + graph.replay() + torch.cuda.synchronize() + + self.assertTrue(torch.isfinite(graph_output).all()) + torch.testing.assert_close(graph_output, eager_output, atol=1e-2, rtol=0) + @torch.no_grad() @unittest.mock.patch( "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None From 8fcea879bfd0a73b25b56817ccb5cfc8235ab9ad Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Thu, 20 Aug 2026 07:04:34 -0700 Subject: [PATCH 2/9] [None][fix] Stabilize Gemma4 FA2 CUDA Graph decode Keep high-head-dimension FlashInfer tensor-core decode plans immutable after CUDA Graph capture while refreshing stable page-table buffers at runtime. Disable split-K for the captured plan, restore Gemma4 CUDA Graph defaults on Hopper, and add long-KV E2B and 12B regressions. Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- .../_torch/attention_backend/flashinfer.py | 25 +++- .../_torch/modeling/test_modeling_gemma4.py | 118 ++++++++++++++---- 2 files changed, 120 insertions(+), 23 deletions(-) diff --git a/tensorrt_llm/_torch/attention_backend/flashinfer.py b/tensorrt_llm/_torch/attention_backend/flashinfer.py index c949137e2be8..5f385a5f8d75 100644 --- a/tensorrt_llm/_torch/attention_backend/flashinfer.py +++ b/tensorrt_llm/_torch/attention_backend/flashinfer.py @@ -233,6 +233,7 @@ class MLAPlanParams: @dataclass(kw_only=True) class FlashInferWrappers: is_planned: bool + cuda_graph_plan_captured: bool = False decode_wrapper: Optional[ flashinfer.BatchDecodeWithPagedKVCacheWrapper] = None prefill_wrapper: Optional[ @@ -1440,7 +1441,13 @@ def _clean_cached_plans(self, *, defer_plan: bool): # corresponding forward pass. So, flush them out here as they won't be relevant for # subsequent forward calls. if plan_params.attention_mask_data is None and plan_params.multi_item_params is None: - self._plan_params_to_wrappers[plan_params].is_planned = False + wrappers = self._plan_params_to_wrappers[plan_params] + # The captured tensor-core kernel owns this plan's launch + # configuration and workspace state. Runtime prepare() still + # refreshes the stable page-table buffers below. + if self.is_cuda_graph and wrappers.cuda_graph_plan_captured: + continue + wrappers.is_planned = False if not defer_plan: self._plan_with_params(plan_params) else: @@ -1895,6 +1902,12 @@ def _plan_with_params(self, plan_params: PlanParams, flashinfer_backend: str = "fa2") -> PlanParams: if not self.needs_plan(plan_params): + if self.is_cuda_graph and torch.cuda.is_current_stream_capturing(): + wrappers = self._plan_params_to_wrappers[plan_params] + if (plan_params.head_dim > 128 + and wrappers.decode_wrapper is not None + and wrappers.decode_wrapper.use_tensor_cores): + wrappers.cuda_graph_plan_captured = True return plan_params if self.is_cuda_graph and torch.cuda.is_current_stream_capturing(): @@ -2015,6 +2028,10 @@ def prefill_plan(): if wrappers.decode_wrapper is None: use_tensor_cores = self._use_tensor_cores(plan_params) + # Gemma4's H256/H512 plans need one immutable tensor-core plan for + # graph capture and replay. The CUDA-core plan is re-planned as KV + # pages change and can mutate state owned by the captured graph. + use_graph_tensor_cores = self.is_cuda_graph and plan_params.head_dim > 128 wrappers.decode_wrapper = \ flashinfer.BatchDecodeWithPagedKVCacheWrapper( @@ -2024,7 +2041,7 @@ def prefill_plan(): paged_kv_indptr_buffer=self.paged_kv_indptr_decode, paged_kv_indices_buffer=self._paged_kv_indices, paged_kv_last_page_len_buffer=self._paged_kv_last_page_len, - use_tensor_cores=use_tensor_cores + use_tensor_cores=use_tensor_cores or use_graph_tensor_cores or flashinfer_backend == "trtllm-gen", backend=flashinfer_backend if flashinfer_backend != "fa2" else @@ -2032,6 +2049,7 @@ def prefill_plan(): 9, 0) else "auto"), ) decode_wrapper = wrappers.decode_wrapper + use_graph_tensor_cores = self.is_cuda_graph and plan_params.head_dim > 128 def decode_plan(): assert decode_wrapper is not None @@ -2068,6 +2086,9 @@ def decode_plan(): 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, + # Split-K can change its launch grid as KV lengths grow. + # Graph replay must retain the grid captured by this plan. + disable_split_kv=use_graph_tensor_cores, ) self._publish_decode_wrapper_kv_lens(decode_wrapper) diff --git a/tests/unittest/_torch/modeling/test_modeling_gemma4.py b/tests/unittest/_torch/modeling/test_modeling_gemma4.py index 26a57427de0c..8ea84d559cd9 100644 --- a/tests/unittest/_torch/modeling/test_modeling_gemma4.py +++ b/tests/unittest/_torch/modeling/test_modeling_gemma4.py @@ -897,6 +897,17 @@ def test_assistant_uses_target_kv_sources(self): "attention_k_eq_v": True, } +# 12B-real-dims: GQA=2 sliding (16/8), GQA=16 full K=V (16/1), +# hd=256/512. +GEMMA4_12B_REAL_DIMS_CONFIG = { + **GEMMA4_E2B_REAL_DIMS_CONFIG, + "num_hidden_layers": 12, + "num_attention_heads": 16, + "num_key_value_heads": 8, + "num_global_key_value_heads": 1, + "attention_k_eq_v": True, +} + # 26B-real-dims: GQA=2 sliding (16/8), GQA=2 full K=V (16/8), hd=256/512. GEMMA4_26B_REAL_DIMS_CONFIG = { **GEMMA4_E2B_REAL_DIMS_CONFIG, @@ -3384,7 +3395,15 @@ def test_cuda_graph_decode_high_gqa(self): @unittest.mock.patch( "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None ) - def _run_cuda_graph_real_headdim(self, config_dict, label=""): + def _run_cuda_graph_real_headdim( + self, + config_dict, + label="", + batch_size=2, + initial_cached=None, + replay_cached=None, + num_blocks=16, + ) -> None: """Helper: CUDA graph decode test with real head_dim configs.""" from tensorrt_llm._torch.attention_backend import ( FlashInferAttention, @@ -3393,17 +3412,30 @@ def _run_cuda_graph_real_headdim(self, config_dict, label=""): from tensorrt_llm._torch.metadata import KVCacheParams config = Gemma4TextConfig(**config_dict) - batch_size = 2 kv_cache_manager = self._get_kv_cache_manager( - config, num_blocks=16, tokens_per_block=32, batch_size=batch_size + config, + num_blocks=num_blocks, + tokens_per_block=32, + batch_size=batch_size, ) self.assertTrue(kv_cache_manager.is_vswa, f"{label}: Expected VSWA manager") request_ids = list(range(batch_size)) - initial_cached = [30, 45] - token_nums = [t + 1 for t in initial_cached] - kv_cache_manager.add_dummy_requests(request_ids, token_nums) + if initial_cached is None: + initial_cached = [30, 45] + self.assertEqual(len(initial_cached), batch_size) + if replay_cached is not None: + self.assertEqual(len(replay_cached), batch_size) + reserved_cached = initial_cached + if replay_cached is not None: + reserved_cached = [ + max(initial, replay) + for initial, replay in zip(initial_cached, replay_cached, strict=True) + ] + token_nums = [t + 1 for t in reserved_cached] + requests = kv_cache_manager.add_dummy_requests(request_ids, token_nums) + self.assertIsNotNone(requests) for i in range(config.num_hidden_layers): buf = kv_cache_manager.get_buffers(i) @@ -3478,21 +3510,6 @@ def _run_cuda_graph_real_headdim(self, config_dict, label=""): ) ) - # --- Reference (eager) --- - ref_metadata = FlashInferAttentionMetadata( - seq_lens=seq_lens, - num_contexts=0, - kv_cache_params=KVCacheParams(use_cache=True, num_cached_tokens_per_seq=initial_cached), - max_num_requests=batch_size, - max_num_tokens=8192, - kv_cache_manager=kv_cache_manager, - request_ids=request_ids, - ) - ref_metadata.prepare() - ref_results = [] - for i in range(num_layers): - ref_results.append(layers[i].forward(gen_qs[i], gen_ks[i], gen_vs[i], ref_metadata)) - # --- CUDA graph --- workspace = torch.empty(320 * 1024 * 1024, dtype=torch.uint8, device="cuda") cg_metadata = FlashInferAttentionMetadata( @@ -3517,8 +3534,35 @@ def _run_cuda_graph_real_headdim(self, config_dict, label=""): with torch.cuda.graph(graph): for i in range(num_layers): cg_results.append(layers[i].forward(gen_qs[i], gen_ks[i], gen_vs[i], cg_metadata)) + + reference_cached = initial_cached + if replay_cached is not None: + reference_cached = replay_cached + cg_metadata.kv_cache_params = KVCacheParams( + use_cache=True, + num_cached_tokens_per_seq=replay_cached, + ) + cg_metadata.prepare() graph.replay() + # --- Reference (eager) --- + ref_metadata = FlashInferAttentionMetadata( + seq_lens=seq_lens, + num_contexts=0, + kv_cache_params=KVCacheParams( + use_cache=True, + num_cached_tokens_per_seq=reference_cached, + ), + max_num_requests=batch_size, + max_num_tokens=8192, + kv_cache_manager=kv_cache_manager, + request_ids=request_ids, + ) + ref_metadata.prepare() + ref_results = [] + for i in range(num_layers): + ref_results.append(layers[i].forward(gen_qs[i], gen_ks[i], gen_vs[i], ref_metadata)) + for i in range(num_layers): torch.testing.assert_close( cg_results[i], @@ -3543,6 +3587,38 @@ def test_cuda_graph_decode_real_headdim(self): """E2B-like: GQA=8, hd=256/512, non-K=V.""" self._run_cuda_graph_real_headdim(deepcopy(GEMMA4_E2B_REAL_DIMS_CONFIG), "E2B") + @torch.no_grad() + @unittest.mock.patch( + "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None + ) + def test_cuda_graph_decode_long_multi_request(self) -> None: + """FA2 CUDA Graph decode matches eager with split-K-sized inputs.""" + batch_size = 8 + self._run_cuda_graph_real_headdim( + deepcopy(GEMMA4_E2B_REAL_DIMS_CONFIG), + "E2B long multi-request", + batch_size=batch_size, + initial_cached=[1024 + 7 * i for i in range(batch_size)], + replay_cached=[30 + i for i in range(batch_size)], + num_blocks=2048, + ) + + @torch.no_grad() + @unittest.mock.patch( + "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None + ) + def test_cuda_graph_decode_long_multi_request_12b_like(self) -> None: + """12B-like FA2 CUDA Graph decode matches eager after capture.""" + batch_size = 1 + self._run_cuda_graph_real_headdim( + deepcopy(GEMMA4_12B_REAL_DIMS_CONFIG), + "12B long multi-request", + batch_size=batch_size, + initial_cached=[1024], + replay_cached=[30], + num_blocks=2048, + ) + @torch.no_grad() @unittest.mock.patch( "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None From ea571f8962cff01dce8c548406cb3836899d07ef Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Thu, 20 Aug 2026 19:36:29 -0700 Subject: [PATCH 3/9] [None][fix] Preserve split-K in FA2 CUDA Graph decode Refresh FlashInfer's split-K schedule outside graph capture when the runtime KV page distribution changes. Keep the captured launch layout and workspace addresses stable, and fail closed if replanning changes the captured plan contract. Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- .../_torch/attention_backend/flashinfer.py | 89 ++++++++-- .../_torch/modeling/test_modeling_gemma4.py | 156 +++++++++++------- 2 files changed, 173 insertions(+), 72 deletions(-) diff --git a/tensorrt_llm/_torch/attention_backend/flashinfer.py b/tensorrt_llm/_torch/attention_backend/flashinfer.py index 5f385a5f8d75..a87b8420fdc0 100644 --- a/tensorrt_llm/_torch/attention_backend/flashinfer.py +++ b/tensorrt_llm/_torch/attention_backend/flashinfer.py @@ -234,6 +234,10 @@ class MLAPlanParams: class FlashInferWrappers: is_planned: bool cuda_graph_plan_captured: bool = False + cuda_graph_plan_info: Optional[tuple[int, ...]] = field(default=None, + repr=False) + decode_plan_num_blocks: Optional[tuple[int, ...]] = field(default=None, + repr=False) decode_wrapper: Optional[ flashinfer.BatchDecodeWithPagedKVCacheWrapper] = None prefill_wrapper: Optional[ @@ -1027,6 +1031,7 @@ def _post_init_with_buffers(self, buffers) -> None: 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._host_paged_kv_last_page_len: Optional[torch.Tensor] = None self._max_num_blocks_per_seq = 0 # VSWA (Variable Sliding Window Attention): models with per-layer @@ -1442,10 +1447,24 @@ def _clean_cached_plans(self, *, defer_plan: bool): # subsequent forward calls. if plan_params.attention_mask_data is None and plan_params.multi_item_params is None: wrappers = self._plan_params_to_wrappers[plan_params] - # The captured tensor-core kernel owns this plan's launch - # configuration and workspace state. Runtime prepare() still - # refreshes the stable page-table buffers below. if self.is_cuda_graph and wrappers.cuda_graph_plan_captured: + if wrappers.cuda_graph_plan_info is None: + continue + num_blocks = tuple(self.num_blocks[self.num_contexts:]) + if wrappers.decode_plan_num_blocks == num_blocks: + continue + wrappers.is_planned = False + self._plan_with_params(plan_params, + synchronize_before_plan=False) + refreshed_plan_info = self._get_fa2_cuda_graph_plan_info( + plan_params, wrappers) + if refreshed_plan_info != wrappers.cuda_graph_plan_info: + raise RuntimeError( + "FlashInfer FA2 CUDA Graph decode plan layout " + f"changed for head_dim={plan_params.head_dim}: " + f"captured={wrappers.cuda_graph_plan_info}, " + f"refreshed={refreshed_plan_info}. The captured " + "graph cannot safely replay this KV distribution.") continue wrappers.is_planned = False if not defer_plan: @@ -1624,6 +1643,7 @@ def _to_int32_tensor(arr: np.ndarray) -> torch.Tensor: paged_kv_last_page_len = _to_int32_tensor(kv_lens_host - (logical_num_blocks - 1) * self.page_size) + self._host_paged_kv_last_page_len = paged_kv_last_page_len self._paged_kv_last_page_len[:paged_kv_last_page_len.size(0)].copy_( paged_kv_last_page_len, non_blocking=True) @@ -1892,6 +1912,28 @@ def _use_tensor_cores(self, plan_params: PlanParams): torch.float8_e4m3fn, torch.float8_e5m2 ] or (plan_params.num_heads // plan_params.num_kv_heads >= 4) + @staticmethod + def _get_fa2_cuda_graph_plan_info( + plan_params: PlanParams, + wrappers: FlashInferWrappers) -> tuple[int, ...]: + decode_wrapper = wrappers.decode_wrapper + if decode_wrapper is None: + raise RuntimeError( + "Missing FlashInfer FA2 CUDA Graph decode wrapper " + f"for head_dim={plan_params.head_dim}.") + plan_info = getattr(decode_wrapper, '_plan_info', None) + if plan_info is None or len(plan_info) != 15: + raise RuntimeError( + "FlashInfer FA2 CUDA Graph tensor-core decode expected a " + f"15-element plan_info for head_dim={plan_params.head_dim}, " + f"got {plan_info}.") + if not plan_info[13] or not plan_info[14]: + raise RuntimeError( + "FlashInfer FA2 CUDA Graph tensor-core decode requires both " + "CUDA Graph and split-K planning for " + f"head_dim={plan_params.head_dim}, got plan_info={plan_info}.") + return tuple(plan_info) + @staticmethod @functools.wraps(flashinfer.BatchPrefillWithPagedKVCacheWrapper) def _page_size_one_prefill_wrapper_builder(*args, **kwargs): @@ -1900,7 +1942,9 @@ def _page_size_one_prefill_wrapper_builder(*args, **kwargs): def _plan_with_params(self, plan_params: PlanParams, - flashinfer_backend: str = "fa2") -> PlanParams: + flashinfer_backend: str = "fa2", + *, + synchronize_before_plan: bool = True) -> PlanParams: if not self.needs_plan(plan_params): if self.is_cuda_graph and torch.cuda.is_current_stream_capturing(): wrappers = self._plan_params_to_wrappers[plan_params] @@ -1908,6 +1952,16 @@ def _plan_with_params(self, and wrappers.decode_wrapper is not None and wrappers.decode_wrapper.use_tensor_cores): wrappers.cuda_graph_plan_captured = True + if wrappers.decode_wrapper._backend == 'fa2': + plan_info = self._get_fa2_cuda_graph_plan_info( + plan_params, wrappers) + if (wrappers.cuda_graph_plan_info is not None + and wrappers.cuda_graph_plan_info != plan_info): + raise RuntimeError( + "FlashInfer FA2 CUDA Graph decode plan changed " + "during capture for " + f"head_dim={plan_params.head_dim}.") + wrappers.cuda_graph_plan_info = plan_info return plan_params if self.is_cuda_graph and torch.cuda.is_current_stream_capturing(): @@ -2028,9 +2082,9 @@ def prefill_plan(): if wrappers.decode_wrapper is None: use_tensor_cores = self._use_tensor_cores(plan_params) - # Gemma4's H256/H512 plans need one immutable tensor-core plan for - # graph capture and replay. The CUDA-core plan is re-planned as KV - # pages change and can mutate state owned by the captured graph. + # Gemma4's H256/H512 plans need a tensor-core wrapper with a stable + # CUDA Graph launch layout. prepare() may refresh its split-K + # schedule in the wrapper's fixed workspace as KV pages change. use_graph_tensor_cores = self.is_cuda_graph and plan_params.head_dim > 128 wrappers.decode_wrapper = \ @@ -2049,7 +2103,6 @@ def prefill_plan(): 9, 0) else "auto"), ) decode_wrapper = wrappers.decode_wrapper - use_graph_tensor_cores = self.is_cuda_graph and plan_params.head_dim > 128 def decode_plan(): assert decode_wrapper is not None @@ -2063,6 +2116,8 @@ def decode_plan(): # its indptr.cpu()/get_seq_lens calls stay free of D2H syncs. paged_kv_indptr = self._host_paged_kv_indptr_decode assert paged_kv_indptr is not None + paged_kv_last_page_len = self._host_paged_kv_last_page_len + assert paged_kv_last_page_len is not None # Persistent, host-built block table: skips flashinfer's # per-request rebuild loop, whose GPU-scalar slice bounds cost # one sync + one scalar D2H per generation request per plan. @@ -2073,7 +2128,7 @@ def decode_plan(): decode_wrapper.plan( paged_kv_indptr[:self.num_generations + 1], self.paged_kv_indices[self.num_context_blocks:], - self.paged_kv_last_page_len[self.num_contexts:], + paged_kv_last_page_len[self.num_contexts:], plan_params.num_heads, plan_params.num_kv_heads, plan_params.head_dim, @@ -2086,14 +2141,20 @@ def decode_plan(): 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, - # Split-K can change its launch grid as KV lengths grow. - # Graph replay must retain the grid captured by this plan. - disable_split_kv=use_graph_tensor_cores, + # FlashInfer pads the FA2 split-K launch grid when + # use_cuda_graph=True. prepare() refreshes the schedule in the + # wrapper's stable workspace and verifies the captured layout. + disable_split_kv=False, ) + wrappers.decode_plan_num_blocks = tuple( + self.num_blocks[self.num_contexts:]) self._publish_decode_wrapper_kv_lens(decode_wrapper) - # Must sync after append_paged_kv_cache and before plan. - torch.cuda.current_stream().synchronize() + # Forward-time planning follows append_paged_kv_cache and retains the + # synchronization. prepare()-time schedule refreshes are ordered on + # the current stream and use host planning metadata, so they skip it. + if synchronize_before_plan: + torch.cuda.current_stream().synchronize() if self.num_contexts > 0: prefill_plan() diff --git a/tests/unittest/_torch/modeling/test_modeling_gemma4.py b/tests/unittest/_torch/modeling/test_modeling_gemma4.py index 8ea84d559cd9..201bb41e4e73 100644 --- a/tests/unittest/_torch/modeling/test_modeling_gemma4.py +++ b/tests/unittest/_torch/modeling/test_modeling_gemma4.py @@ -3397,12 +3397,13 @@ def test_cuda_graph_decode_high_gqa(self): ) def _run_cuda_graph_real_headdim( self, - config_dict, - label="", - batch_size=2, - initial_cached=None, - replay_cached=None, - num_blocks=16, + config_dict: dict, + label: str = "", + batch_size: int = 2, + initial_cached: list[int] | None = None, + replay_cached_steps: list[list[int]] | None = None, + num_blocks: int = 16, + expect_split_kv: bool = False, ) -> None: """Helper: CUDA graph decode test with real head_dim configs.""" from tensorrt_llm._torch.attention_backend import ( @@ -3425,14 +3426,14 @@ def _run_cuda_graph_real_headdim( if initial_cached is None: initial_cached = [30, 45] self.assertEqual(len(initial_cached), batch_size) - if replay_cached is not None: + if replay_cached_steps is None: + replay_cached_steps = [initial_cached] + for replay_cached in replay_cached_steps: self.assertEqual(len(replay_cached), batch_size) - reserved_cached = initial_cached - if replay_cached is not None: - reserved_cached = [ - max(initial, replay) - for initial, replay in zip(initial_cached, replay_cached, strict=True) - ] + reserved_cached = [ + max(cached_tokens) + for cached_tokens in zip(initial_cached, *replay_cached_steps, strict=True) + ] token_nums = [t + 1 for t in reserved_cached] requests = kv_cache_manager.add_dummy_requests(request_ids, token_nums) self.assertIsNotNone(requests) @@ -3535,47 +3536,76 @@ def _run_cuda_graph_real_headdim( for i in range(num_layers): cg_results.append(layers[i].forward(gen_qs[i], gen_ks[i], gen_vs[i], cg_metadata)) - reference_cached = initial_cached - if replay_cached is not None: - reference_cached = replay_cached + captured_split_kv_plans = {} + if expect_split_kv and torch.cuda.get_device_capability() == (9, 0): + split_kv_head_dims = set() + for plan_params, wrappers in cg_metadata._plan_params_to_wrappers.items(): + decode_wrapper = wrappers.decode_wrapper + if decode_wrapper is None or decode_wrapper._backend != "fa2": + continue + self.assertTrue(wrappers.cuda_graph_plan_captured) + self.assertEqual(len(decode_wrapper._plan_info), 15) + self.assertTrue( + decode_wrapper._plan_info[14], + f"{label}: FA2 hd={plan_params.head_dim} did not enable split-K", + ) + self.assertEqual(wrappers.cuda_graph_plan_info, tuple(decode_wrapper._plan_info)) + captured_split_kv_plans[plan_params] = ( + tuple(decode_wrapper._plan_info), + decode_wrapper._int_workspace_buffer.data_ptr(), + ) + split_kv_head_dims.add(plan_params.head_dim) + self.assertEqual(split_kv_head_dims, {256, 512}) + + for replay_step, reference_cached in enumerate(replay_cached_steps): cg_metadata.kv_cache_params = KVCacheParams( - use_cache=True, - num_cached_tokens_per_seq=replay_cached, + use_cache=True, num_cached_tokens_per_seq=reference_cached ) cg_metadata.prepare() - graph.replay() - # --- Reference (eager) --- - ref_metadata = FlashInferAttentionMetadata( - seq_lens=seq_lens, - num_contexts=0, - kv_cache_params=KVCacheParams( - use_cache=True, - num_cached_tokens_per_seq=reference_cached, - ), - max_num_requests=batch_size, - max_num_tokens=8192, - kv_cache_manager=kv_cache_manager, - request_ids=request_ids, - ) - ref_metadata.prepare() - ref_results = [] - for i in range(num_layers): - ref_results.append(layers[i].forward(gen_qs[i], gen_ks[i], gen_vs[i], ref_metadata)) + for plan_params, (captured_plan_info, workspace_ptr) in captured_split_kv_plans.items(): + wrappers = cg_metadata._plan_params_to_wrappers[plan_params] + decode_wrapper = wrappers.decode_wrapper + self.assertEqual(tuple(decode_wrapper._plan_info), captured_plan_info) + self.assertEqual(decode_wrapper._int_workspace_buffer.data_ptr(), workspace_ptr) + self.assertEqual( + wrappers.decode_plan_num_blocks, + tuple(cg_metadata.num_blocks[cg_metadata.num_contexts :]), + ) - for i in range(num_layers): - torch.testing.assert_close( - cg_results[i], - ref_results[i], - atol=1e-2, - rtol=0, - msg=( - f"{label} Layer {i} ({layer_types[i]}, " - f"hd={layers_info[i]['head_dim']}, " - f"kv={layers_info[i]['num_kv_heads']}): " - f"CUDA graph diverges from eager" + graph.replay() + + # --- Reference (eager) --- + ref_metadata = FlashInferAttentionMetadata( + seq_lens=seq_lens, + num_contexts=0, + kv_cache_params=KVCacheParams( + use_cache=True, + num_cached_tokens_per_seq=reference_cached, ), + max_num_requests=batch_size, + max_num_tokens=8192, + kv_cache_manager=kv_cache_manager, + request_ids=request_ids, ) + ref_metadata.prepare() + ref_results = [] + for i in range(num_layers): + ref_results.append(layers[i].forward(gen_qs[i], gen_ks[i], gen_vs[i], ref_metadata)) + + for i in range(num_layers): + torch.testing.assert_close( + cg_results[i], + ref_results[i], + atol=1e-2, + rtol=0, + msg=( + f"{label} replay {replay_step}, Layer {i} ({layer_types[i]}, " + f"hd={layers_info[i]['head_dim']}, " + f"kv={layers_info[i]['num_kv_heads']}): " + f"CUDA graph diverges from eager" + ), + ) kv_cache_manager.shutdown() @@ -3592,15 +3622,20 @@ def test_cuda_graph_decode_real_headdim(self): "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None ) def test_cuda_graph_decode_long_multi_request(self) -> None: - """FA2 CUDA Graph decode matches eager with split-K-sized inputs.""" + """E2B-like split-K schedules refresh across graph replays.""" batch_size = 8 self._run_cuda_graph_real_headdim( deepcopy(GEMMA4_E2B_REAL_DIMS_CONFIG), - "E2B long multi-request", + "E2B split-K schedule refresh", batch_size=batch_size, - initial_cached=[1024 + 7 * i for i in range(batch_size)], - replay_cached=[30 + i for i in range(batch_size)], - num_blocks=2048, + initial_cached=[4095] + [31] * (batch_size - 1), + replay_cached_steps=[ + [511] * batch_size, + [31] * (batch_size - 1) + [4095], + [30, 31, 32, 33, 1023, 1024, 1025, 2048], + ], + num_blocks=8192, + expect_split_kv=True, ) @torch.no_grad() @@ -3608,15 +3643,20 @@ def test_cuda_graph_decode_long_multi_request(self) -> None: "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None ) def test_cuda_graph_decode_long_multi_request_12b_like(self) -> None: - """12B-like FA2 CUDA Graph decode matches eager after capture.""" - batch_size = 1 + """12B-like split-K schedules refresh across graph replays.""" + batch_size = 8 self._run_cuda_graph_real_headdim( deepcopy(GEMMA4_12B_REAL_DIMS_CONFIG), - "12B long multi-request", + "12B split-K schedule refresh", batch_size=batch_size, - initial_cached=[1024], - replay_cached=[30], - num_blocks=2048, + initial_cached=[4095] + [31] * (batch_size - 1), + replay_cached_steps=[ + [511] * batch_size, + [31] * (batch_size - 1) + [4095], + [30, 31, 32, 33, 1023, 1024, 1025, 2048], + ], + num_blocks=8192, + expect_split_kv=True, ) @torch.no_grad() From 91d438e1baca817c3cf2a28c540d522ced992533 Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Thu, 20 Aug 2026 20:57:50 -0700 Subject: [PATCH 4/9] [None][fix] Simplify Gemma4 FA2 graph schedule refresh Rely on FlashInfer's pinned CUDA Graph plan contract instead of duplicating private plan_info validation. Refresh only when FA2 page-count distributions change, preserve the conservative KV bound, and keep the original synchronization required by the pinned planning workspace. Consolidate Hopper coverage and reduce oversized test allocations. Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- .../_torch/attention_backend/flashinfer.py | 97 ++------ .../test_lists/test-db/l0_h100.yml | 1 + .../_torch/modeling/test_modeling_gemma4.py | 229 ++++-------------- 3 files changed, 68 insertions(+), 259 deletions(-) diff --git a/tensorrt_llm/_torch/attention_backend/flashinfer.py b/tensorrt_llm/_torch/attention_backend/flashinfer.py index a87b8420fdc0..d2a369c219a2 100644 --- a/tensorrt_llm/_torch/attention_backend/flashinfer.py +++ b/tensorrt_llm/_torch/attention_backend/flashinfer.py @@ -233,11 +233,8 @@ class MLAPlanParams: @dataclass(kw_only=True) class FlashInferWrappers: is_planned: bool - cuda_graph_plan_captured: bool = False - cuda_graph_plan_info: Optional[tuple[int, ...]] = field(default=None, - repr=False) - decode_plan_num_blocks: Optional[tuple[int, ...]] = field(default=None, - repr=False) + fa2_plan_num_blocks: Optional[tuple[int, ...]] = field(default=None, + repr=False) decode_wrapper: Optional[ flashinfer.BatchDecodeWithPagedKVCacheWrapper] = None prefill_wrapper: Optional[ @@ -1447,24 +1444,12 @@ def _clean_cached_plans(self, *, defer_plan: bool): # subsequent forward calls. if plan_params.attention_mask_data is None and plan_params.multi_item_params is None: wrappers = self._plan_params_to_wrappers[plan_params] - if self.is_cuda_graph and wrappers.cuda_graph_plan_captured: - if wrappers.cuda_graph_plan_info is None: - continue + if wrappers.fa2_plan_num_blocks is not None: num_blocks = tuple(self.num_blocks[self.num_contexts:]) - if wrappers.decode_plan_num_blocks == num_blocks: + if not num_blocks or wrappers.fa2_plan_num_blocks == num_blocks: continue wrappers.is_planned = False - self._plan_with_params(plan_params, - synchronize_before_plan=False) - refreshed_plan_info = self._get_fa2_cuda_graph_plan_info( - plan_params, wrappers) - if refreshed_plan_info != wrappers.cuda_graph_plan_info: - raise RuntimeError( - "FlashInfer FA2 CUDA Graph decode plan layout " - f"changed for head_dim={plan_params.head_dim}: " - f"captured={wrappers.cuda_graph_plan_info}, " - f"refreshed={refreshed_plan_info}. The captured " - "graph cannot safely replay this KV distribution.") + self._plan_with_params(plan_params) continue wrappers.is_planned = False if not defer_plan: @@ -1643,7 +1628,6 @@ def _to_int32_tensor(arr: np.ndarray) -> torch.Tensor: paged_kv_last_page_len = _to_int32_tensor(kv_lens_host - (logical_num_blocks - 1) * self.page_size) - self._host_paged_kv_last_page_len = paged_kv_last_page_len self._paged_kv_last_page_len[:paged_kv_last_page_len.size(0)].copy_( paged_kv_last_page_len, non_blocking=True) @@ -1912,28 +1896,6 @@ def _use_tensor_cores(self, plan_params: PlanParams): torch.float8_e4m3fn, torch.float8_e5m2 ] or (plan_params.num_heads // plan_params.num_kv_heads >= 4) - @staticmethod - def _get_fa2_cuda_graph_plan_info( - plan_params: PlanParams, - wrappers: FlashInferWrappers) -> tuple[int, ...]: - decode_wrapper = wrappers.decode_wrapper - if decode_wrapper is None: - raise RuntimeError( - "Missing FlashInfer FA2 CUDA Graph decode wrapper " - f"for head_dim={plan_params.head_dim}.") - plan_info = getattr(decode_wrapper, '_plan_info', None) - if plan_info is None or len(plan_info) != 15: - raise RuntimeError( - "FlashInfer FA2 CUDA Graph tensor-core decode expected a " - f"15-element plan_info for head_dim={plan_params.head_dim}, " - f"got {plan_info}.") - if not plan_info[13] or not plan_info[14]: - raise RuntimeError( - "FlashInfer FA2 CUDA Graph tensor-core decode requires both " - "CUDA Graph and split-K planning for " - f"head_dim={plan_params.head_dim}, got plan_info={plan_info}.") - return tuple(plan_info) - @staticmethod @functools.wraps(flashinfer.BatchPrefillWithPagedKVCacheWrapper) def _page_size_one_prefill_wrapper_builder(*args, **kwargs): @@ -1942,26 +1904,8 @@ def _page_size_one_prefill_wrapper_builder(*args, **kwargs): def _plan_with_params(self, plan_params: PlanParams, - flashinfer_backend: str = "fa2", - *, - synchronize_before_plan: bool = True) -> PlanParams: + flashinfer_backend: str = "fa2") -> PlanParams: if not self.needs_plan(plan_params): - if self.is_cuda_graph and torch.cuda.is_current_stream_capturing(): - wrappers = self._plan_params_to_wrappers[plan_params] - if (plan_params.head_dim > 128 - and wrappers.decode_wrapper is not None - and wrappers.decode_wrapper.use_tensor_cores): - wrappers.cuda_graph_plan_captured = True - if wrappers.decode_wrapper._backend == 'fa2': - plan_info = self._get_fa2_cuda_graph_plan_info( - plan_params, wrappers) - if (wrappers.cuda_graph_plan_info is not None - and wrappers.cuda_graph_plan_info != plan_info): - raise RuntimeError( - "FlashInfer FA2 CUDA Graph decode plan changed " - "during capture for " - f"head_dim={plan_params.head_dim}.") - wrappers.cuda_graph_plan_info = plan_info return plan_params if self.is_cuda_graph and torch.cuda.is_current_stream_capturing(): @@ -2080,12 +2024,12 @@ def prefill_plan(): custom_mask=plan_params.attention_mask_data, ) + use_graph_tensor_cores = self.is_cuda_graph and plan_params.head_dim > 128 if wrappers.decode_wrapper is None: use_tensor_cores = self._use_tensor_cores(plan_params) # Gemma4's H256/H512 plans need a tensor-core wrapper with a stable # CUDA Graph launch layout. prepare() may refresh its split-K # schedule in the wrapper's fixed workspace as KV pages change. - use_graph_tensor_cores = self.is_cuda_graph and plan_params.head_dim > 128 wrappers.decode_wrapper = \ flashinfer.BatchDecodeWithPagedKVCacheWrapper( @@ -2099,7 +2043,7 @@ def prefill_plan(): or flashinfer_backend == "trtllm-gen", backend=flashinfer_backend if flashinfer_backend != "fa2" else - ("fa2" if torch.cuda.get_device_capability(0) == ( + ("fa2" if torch.cuda.get_device_capability() == ( 9, 0) else "auto"), ) decode_wrapper = wrappers.decode_wrapper @@ -2116,8 +2060,6 @@ def decode_plan(): # its indptr.cpu()/get_seq_lens calls stay free of D2H syncs. paged_kv_indptr = self._host_paged_kv_indptr_decode assert paged_kv_indptr is not None - paged_kv_last_page_len = self._host_paged_kv_last_page_len - assert paged_kv_last_page_len is not None # Persistent, host-built block table: skips flashinfer's # per-request rebuild loop, whose GPU-scalar slice bounds cost # one sync + one scalar D2H per generation request per plan. @@ -2125,10 +2067,13 @@ def decode_plan(): if decode_wrapper._backend == 'trtllm-gen': block_tables = self._build_decode_block_tables( plan_params, wrappers) + planned_max_kv_len = (decode_wrapper._max_kv_len + if wrappers.fa2_plan_num_blocks is not None + else None) decode_wrapper.plan( paged_kv_indptr[:self.num_generations + 1], self.paged_kv_indices[self.num_context_blocks:], - paged_kv_last_page_len[self.num_contexts:], + self.paged_kv_last_page_len[self.num_contexts:], plan_params.num_heads, plan_params.num_kv_heads, plan_params.head_dim, @@ -2141,20 +2086,18 @@ def decode_plan(): 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, - # FlashInfer pads the FA2 split-K launch grid when - # use_cuda_graph=True. prepare() refreshes the schedule in the - # wrapper's stable workspace and verifies the captured layout. disable_split_kv=False, ) - wrappers.decode_plan_num_blocks = tuple( - self.num_blocks[self.num_contexts:]) + if planned_max_kv_len is not None: + decode_wrapper._max_kv_len = max(planned_max_kv_len, + decode_wrapper._max_kv_len) + if use_graph_tensor_cores and decode_wrapper._backend == 'fa2': + wrappers.fa2_plan_num_blocks = tuple( + self.num_blocks[self.num_contexts:]) self._publish_decode_wrapper_kv_lens(decode_wrapper) - # Forward-time planning follows append_paged_kv_cache and retains the - # synchronization. prepare()-time schedule refreshes are ordered on - # the current stream and use host planning metadata, so they skip it. - if synchronize_before_plan: - torch.cuda.current_stream().synchronize() + # Must sync after append_paged_kv_cache and before plan. + torch.cuda.current_stream().synchronize() if self.num_contexts > 0: prefill_plan() diff --git a/tests/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index 79d77502baf4..36c7673cf565 100644 --- a/tests/integration/test_lists/test-db/l0_h100.yml +++ b/tests/integration/test_lists/test-db/l0_h100.yml @@ -68,6 +68,7 @@ l0_h100: - unittest/_torch/modeling -k "modeling_mixtral" - unittest/_torch/modeling -k "modeling_gemma3" - unittest/_torch/modeling -k "modeling_gpt_oss" + - unittest/_torch/modeling/test_modeling_gemma4.py::TestGemma4CUDAGraph::test_cuda_graph_split_kv_schedule_refresh - unittest/_torch/modeling -k "modeling_whisper" # CPU-only log-mel parity - unittest/_torch/modeling/test_modeling_nemotron_h.py::test_nemotron_h_sanity # Real-weight Nano CG/overlap and chunked-prefill path smoke (MoE L0 cannot diff --git a/tests/unittest/_torch/modeling/test_modeling_gemma4.py b/tests/unittest/_torch/modeling/test_modeling_gemma4.py index 201bb41e4e73..0fafdd2ca67e 100644 --- a/tests/unittest/_torch/modeling/test_modeling_gemma4.py +++ b/tests/unittest/_torch/modeling/test_modeling_gemma4.py @@ -924,6 +924,7 @@ def _build_gemma4_kv_cache_manager( num_blocks=4, tokens_per_block=32, batch_size=1, + enable_swa_eviction: bool = False, ): """Create KVCacheManagerV2 supporting Gemma4 per-layer head_dim / kv_heads. @@ -981,7 +982,9 @@ def _build_gemma4_kv_cache_manager( if not needs_vswa: needs_vswa = isinstance(num_kv_heads, list) and len(set(num_kv_heads)) > 1 if needs_vswa and sliding_window: - swa_window = min(sliding_window, max_seq_len - 1) + swa_window = ( + min(sliding_window, max_seq_len - 1) if enable_swa_eviction else max_seq_len - 1 + ) max_attn_window = [ swa_window if lt == "sliding_attention" else max_seq_len for lt in layer_types ] @@ -2075,21 +2078,8 @@ def test_vswa_pool_cache_not_aliased(self): @unittest.mock.patch( "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None ) - def test_vswa_no_eviction_with_long_sequence(self): - """VSWA: sliding pool must not evict blocks when max_attention_window - uses max_seq_len - 1 (the fix for page index OOB). - - Root cause: when _util.py used the model's sliding_window (e.g. 512) - as max_attention_window for sliding layers, V2 would evict old blocks - when kv_lens exceeded the window. But FlashInfer's prepare() computes - num_blocks from the FULL kv_lens, so the page indices for evicted - blocks become stale → illegal memory access. - - The fix uses max_seq_len - 1 instead of sliding_window, preventing - eviction while keeping is_vswa=True. This test verifies that with - the fix, a sequence longer than sliding_window still has all its - blocks allocated (no eviction) and page indices are within bounds. - """ + def test_vswa_evicted_page_indices_are_sanitized(self) -> None: + """FlashInfer metadata replaces evicted SWA page markers.""" from tensorrt_llm._torch.attention_backend.utils import get_attention_backend from tensorrt_llm._torch.metadata import KVCacheParams @@ -2098,44 +2088,45 @@ def test_vswa_no_eviction_with_long_sequence(self): config_dict["sliding_window"] = 64 config = Gemma4TextConfig(**config_dict) - # num_blocks=4 → max_seq_len = 4*128 = 512, much larger than - # sliding_window=64. With the fix, max_attention_window for sliding - # layers = 511 (max_seq_len - 1), so V2 won't evict. - kv_cache_manager = self._get_kv_cache_manager(config, num_blocks=4) + kv_cache_manager = self._get_kv_cache_manager( + config, num_blocks=4, enable_swa_eviction=True + ) - # Allocate a request with tokens > sliding_window + # Allocate a generation request longer than the sliding window. request_ids = [1] - token_nums = [128] # 128 tokens >> sliding_window (64) - kv_cache_manager.add_dummy_requests(request_ids, token_nums) + cached_tokens = 126 + token_nums = [cached_tokens + 1] + kv_cache_manager.add_dummy_requests(request_ids, token_nums, is_gen=True) + + num_blocks = (token_nums[0] + kv_cache_manager.tokens_per_block - 1) // ( + kv_cache_manager.tokens_per_block + ) + raw_indices = kv_cache_manager.get_batch_cache_indices_flat( + request_ids, [num_blocks], layer_idx=0 + ) + self.assertIn(BAD_PAGE_INDEX, raw_indices.tolist()) metadata_cls = get_attention_backend("FLASHINFER").Metadata metadata = metadata_cls( - seq_lens=torch.tensor([128], dtype=torch.int), - num_contexts=1, + seq_lens=torch.ones(1, dtype=torch.int), + num_contexts=0, kv_cache_params=KVCacheParams( use_cache=True, - num_cached_tokens_per_seq=[0], + num_cached_tokens_per_seq=[cached_tokens], ), max_num_requests=1, max_num_tokens=8192, kv_cache_manager=kv_cache_manager, request_ids=request_ids, - prompt_lens=[128], ) with torch.inference_mode(): metadata.prepare() - # num_blocks should be based on full kv_lens (128 tokens), - # not clamped to sliding_window (64 tokens). - expected_blocks = ( - 128 + kv_cache_manager.tokens_per_block - 1 - ) // kv_cache_manager.tokens_per_block - self.assertEqual( - metadata.num_blocks[0], - expected_blocks, - f"num_blocks should be {expected_blocks} (from full kv_lens=128), " - f"not clamped to sliding_window={config_dict['sliding_window']}", + self.assertEqual(metadata.num_blocks[0], num_blocks) + self.assertNotIn( + BAD_PAGE_INDEX, + metadata.get_paged_kv_indices_for_layer(0).cpu().tolist(), ) # Page indices must be within bounds for EVERY layer @@ -3401,7 +3392,7 @@ def _run_cuda_graph_real_headdim( label: str = "", batch_size: int = 2, initial_cached: list[int] | None = None, - replay_cached_steps: list[list[int]] | None = None, + replay_cached: list[int] | None = None, num_blocks: int = 16, expect_split_kv: bool = False, ) -> None: @@ -3426,16 +3417,15 @@ def _run_cuda_graph_real_headdim( if initial_cached is None: initial_cached = [30, 45] self.assertEqual(len(initial_cached), batch_size) - if replay_cached_steps is None: - replay_cached_steps = [initial_cached] - for replay_cached in replay_cached_steps: - self.assertEqual(len(replay_cached), batch_size) + if replay_cached is None: + replay_cached = initial_cached + self.assertEqual(len(replay_cached), batch_size) reserved_cached = [ - max(cached_tokens) - for cached_tokens in zip(initial_cached, *replay_cached_steps, strict=True) + max(initial, replay) + for initial, replay in zip(initial_cached, replay_cached, strict=True) ] token_nums = [t + 1 for t in reserved_cached] - requests = kv_cache_manager.add_dummy_requests(request_ids, token_nums) + requests = kv_cache_manager.add_dummy_requests(request_ids, token_nums, is_gen=True) self.assertIsNotNone(requests) for i in range(config.num_hidden_layers): @@ -3536,43 +3526,28 @@ def _run_cuda_graph_real_headdim( for i in range(num_layers): cg_results.append(layers[i].forward(gen_qs[i], gen_ks[i], gen_vs[i], cg_metadata)) - captured_split_kv_plans = {} - if expect_split_kv and torch.cuda.get_device_capability() == (9, 0): + if expect_split_kv: split_kv_head_dims = set() for plan_params, wrappers in cg_metadata._plan_params_to_wrappers.items(): decode_wrapper = wrappers.decode_wrapper if decode_wrapper is None or decode_wrapper._backend != "fa2": continue - self.assertTrue(wrappers.cuda_graph_plan_captured) - self.assertEqual(len(decode_wrapper._plan_info), 15) self.assertTrue( - decode_wrapper._plan_info[14], + decode_wrapper._plan_info[-1], f"{label}: FA2 hd={plan_params.head_dim} did not enable split-K", ) - self.assertEqual(wrappers.cuda_graph_plan_info, tuple(decode_wrapper._plan_info)) - captured_split_kv_plans[plan_params] = ( - tuple(decode_wrapper._plan_info), - decode_wrapper._int_workspace_buffer.data_ptr(), - ) split_kv_head_dims.add(plan_params.head_dim) self.assertEqual(split_kv_head_dims, {256, 512}) + replay_cached_steps = [replay_cached] + if expect_split_kv: + replay_cached_steps.append([cached + 1 for cached in replay_cached]) for replay_step, reference_cached in enumerate(replay_cached_steps): cg_metadata.kv_cache_params = KVCacheParams( use_cache=True, num_cached_tokens_per_seq=reference_cached ) cg_metadata.prepare() - for plan_params, (captured_plan_info, workspace_ptr) in captured_split_kv_plans.items(): - wrappers = cg_metadata._plan_params_to_wrappers[plan_params] - decode_wrapper = wrappers.decode_wrapper - self.assertEqual(tuple(decode_wrapper._plan_info), captured_plan_info) - self.assertEqual(decode_wrapper._int_workspace_buffer.data_ptr(), workspace_ptr) - self.assertEqual( - wrappers.decode_plan_num_blocks, - tuple(cg_metadata.num_blocks[cg_metadata.num_contexts :]), - ) - graph.replay() # --- Reference (eager) --- @@ -3621,41 +3596,20 @@ def test_cuda_graph_decode_real_headdim(self): @unittest.mock.patch( "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None ) - def test_cuda_graph_decode_long_multi_request(self) -> None: - """E2B-like split-K schedules refresh across graph replays.""" - batch_size = 8 - self._run_cuda_graph_real_headdim( - deepcopy(GEMMA4_E2B_REAL_DIMS_CONFIG), - "E2B split-K schedule refresh", - batch_size=batch_size, - initial_cached=[4095] + [31] * (batch_size - 1), - replay_cached_steps=[ - [511] * batch_size, - [31] * (batch_size - 1) + [4095], - [30, 31, 32, 33, 1023, 1024, 1025, 2048], - ], - num_blocks=8192, - expect_split_kv=True, - ) - - @torch.no_grad() - @unittest.mock.patch( - "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None + @unittest.skipUnless( + torch.cuda.is_available() and torch.cuda.get_device_capability() == (9, 0), + "FA2 split-K schedule refresh is Hopper-specific", ) - def test_cuda_graph_decode_long_multi_request_12b_like(self) -> None: - """12B-like split-K schedules refresh across graph replays.""" + def test_cuda_graph_split_kv_schedule_refresh(self) -> None: + """FA2 split-K graphs refresh schedules for new KV distributions.""" batch_size = 8 self._run_cuda_graph_real_headdim( deepcopy(GEMMA4_12B_REAL_DIMS_CONFIG), "12B split-K schedule refresh", batch_size=batch_size, initial_cached=[4095] + [31] * (batch_size - 1), - replay_cached_steps=[ - [511] * batch_size, - [31] * (batch_size - 1) + [4095], - [30, 31, 32, 33, 1023, 1024, 1025, 2048], - ], - num_blocks=8192, + replay_cached=[510] * batch_size, + num_blocks=512, expect_split_kv=True, ) @@ -3675,95 +3629,6 @@ def test_cuda_graph_decode_26b_like(self): """26B-like: GQA=2, K=V, hd=256/512.""" self._run_cuda_graph_real_headdim(deepcopy(GEMMA4_26B_REAL_DIMS_CONFIG), "26B") - @torch.no_grad() - @unittest.mock.patch( - "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None - ) - def test_cuda_graph_decode_with_evicted_swa_pages(self): - """FA2 CUDA graph decode safely ignores evicted SWA pages.""" - config_dict = deepcopy(GEMMA4_26B_REAL_DIMS_CONFIG) - config_dict["sliding_window"] = 1024 - config_dict["max_position_embeddings"] = 4096 - config = Gemma4TextConfig(**config_dict) - - tokens_per_block = 32 - cached_tokens = 2046 - kv_cache_manager = self._get_kv_cache_manager( - config, - num_blocks=80, - tokens_per_block=tokens_per_block, - batch_size=1, - ) - self.addCleanup(kv_cache_manager.shutdown) - - requests = kv_cache_manager.add_dummy_requests([0], [cached_tokens + 1], is_gen=True) - self.assertIsNotNone(requests) - torch.nn.init.normal_(kv_cache_manager.get_buffers(0)) - - num_blocks = (cached_tokens + 1 + tokens_per_block - 1) // tokens_per_block - raw_indices = kv_cache_manager.get_batch_cache_indices_flat([0], [num_blocks], layer_idx=0) - self.assertIn(BAD_PAGE_INDEX, raw_indices.tolist()) - - metadata = FlashInferAttentionMetadata( - seq_lens=torch.ones(1, dtype=torch.int), - num_contexts=0, - is_cuda_graph=True, - kv_cache_params=KVCacheParams( - use_cache=True, - num_cached_tokens_per_seq=[cached_tokens], - ), - workspace_buffer=torch.empty( - _FLASHINFER_WORKSPACE_BYTES, dtype=torch.uint8, device="cuda" - ), - max_num_requests=1, - max_num_tokens=4096, - kv_cache_manager=kv_cache_manager, - request_ids=[0], - ) - metadata.prepare() - sanitized_indices = metadata.get_paged_kv_indices_for_layer(0) - self.assertNotIn(BAD_PAGE_INDEX, sanitized_indices.cpu().tolist()) - - layer = FlashInferAttention( - layer_idx=0, - num_heads=config.num_attention_heads, - head_dim=config.head_dim, - num_kv_heads=config.num_key_value_heads, - q_scaling=1.0 / math.sqrt(config.head_dim), - flashinfer_backend="fa2", - ) - query = torch.randn( - 1, - config.num_attention_heads * config.head_dim, - dtype=config.torch_dtype, - device="cuda", - ) - - eager_output = None - for _ in range(2): - eager_output = layer.forward( - query, - None, - None, - metadata, - attention_window_size=config.sliding_window, - ) - - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - graph_output = layer.forward( - query, - None, - None, - metadata, - attention_window_size=config.sliding_window, - ) - graph.replay() - torch.cuda.synchronize() - - self.assertTrue(torch.isfinite(graph_output).all()) - torch.testing.assert_close(graph_output, eager_output, atol=1e-2, rtol=0) - @torch.no_grad() @unittest.mock.patch( "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None From 911015f83e1eed6e12a3c7aa3f017a9ca9bd76ab Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Thu, 20 Aug 2026 21:30:06 -0700 Subject: [PATCH 5/9] [None][test] Align Gemma4 H100 test selection Use the modeling directory filter format shared by the other H100 model entries so all Gemma4 modeling tests are selected consistently. Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- tests/integration/test_lists/test-db/l0_h100.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/integration/test_lists/test-db/l0_h100.yml b/tests/integration/test_lists/test-db/l0_h100.yml index 36c7673cf565..e62cbbc3f9bc 100644 --- a/tests/integration/test_lists/test-db/l0_h100.yml +++ b/tests/integration/test_lists/test-db/l0_h100.yml @@ -68,7 +68,7 @@ l0_h100: - unittest/_torch/modeling -k "modeling_mixtral" - unittest/_torch/modeling -k "modeling_gemma3" - unittest/_torch/modeling -k "modeling_gpt_oss" - - unittest/_torch/modeling/test_modeling_gemma4.py::TestGemma4CUDAGraph::test_cuda_graph_split_kv_schedule_refresh + - unittest/_torch/modeling -k "modeling_gemma4" - unittest/_torch/modeling -k "modeling_whisper" # CPU-only log-mel parity - unittest/_torch/modeling/test_modeling_nemotron_h.py::test_nemotron_h_sanity # Real-weight Nano CG/overlap and chunked-prefill path smoke (MoE L0 cannot From bc97ea95fe37c6d98450dce577e200493efb87d6 Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Thu, 20 Aug 2026 21:45:22 -0700 Subject: [PATCH 6/9] [None][fix] Handle Gemma4 graph mode transitions Clear the cached FA2 decode signature before a prefill-only plan and mark trtllm-gen-only CUDA Graph tests as SM100f-only. This keeps the H100 Gemma4 modeling filter from collecting unsupported kernels while preserving Hopper FA2 coverage. Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- tensorrt_llm/_torch/attention_backend/flashinfer.py | 11 +++++++---- .../unittest/_torch/modeling/test_modeling_gemma4.py | 10 ++++++++-- 2 files changed, 15 insertions(+), 6 deletions(-) diff --git a/tensorrt_llm/_torch/attention_backend/flashinfer.py b/tensorrt_llm/_torch/attention_backend/flashinfer.py index d2a369c219a2..737f8769ad97 100644 --- a/tensorrt_llm/_torch/attention_backend/flashinfer.py +++ b/tensorrt_llm/_torch/attention_backend/flashinfer.py @@ -1446,11 +1446,14 @@ def _clean_cached_plans(self, *, defer_plan: bool): wrappers = self._plan_params_to_wrappers[plan_params] if wrappers.fa2_plan_num_blocks is not None: num_blocks = tuple(self.num_blocks[self.num_contexts:]) - if not num_blocks or wrappers.fa2_plan_num_blocks == num_blocks: + if not num_blocks: + wrappers.fa2_plan_num_blocks = None + elif wrappers.fa2_plan_num_blocks == num_blocks: + continue + else: + wrappers.is_planned = False + self._plan_with_params(plan_params) continue - wrappers.is_planned = False - self._plan_with_params(plan_params) - continue wrappers.is_planned = False if not defer_plan: self._plan_with_params(plan_params) diff --git a/tests/unittest/_torch/modeling/test_modeling_gemma4.py b/tests/unittest/_torch/modeling/test_modeling_gemma4.py index 0fafdd2ca67e..d464882a7d14 100644 --- a/tests/unittest/_torch/modeling/test_modeling_gemma4.py +++ b/tests/unittest/_torch/modeling/test_modeling_gemma4.py @@ -2662,6 +2662,7 @@ def _expected_decode_block_table( source_offset += page_count return expected + @unittest.skipUnless(is_sm_100f(), "trtllm-gen attention requires SM100f") @torch.no_grad() @unittest.mock.patch( "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None @@ -2703,6 +2704,7 @@ def test_shared_kv_draft_view(self) -> None: rtol=0, ) + @unittest.skipUnless(is_sm_100f(), "trtllm-gen attention requires SM100f") @torch.no_grad() @unittest.mock.patch( "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None @@ -2740,6 +2742,7 @@ def test_cuda_graph_trtllm_gen_block_table_transitions(self) -> None: self.assertEqual(wrappers.decode_block_table_active_rows, len(new_page_counts)) self.assertEqual(wrappers.decode_block_table_active_width, max(new_page_counts)) + @unittest.skipUnless(is_sm_100f(), "trtllm-gen attention requires SM100f") @torch.no_grad() @unittest.mock.patch( "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None @@ -2790,6 +2793,7 @@ def test_cuda_graph_trtllm_gen_host_table_growth_keeps_device_pointer(self) -> N rtol=0, ) + @unittest.skipUnless(is_sm_100f(), "trtllm-gen attention requires SM100f") @torch.no_grad() @unittest.mock.patch( "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None @@ -3215,11 +3219,12 @@ def test_cuda_graph_multi_step_decode(self): kv_cache_manager.shutdown() + @unittest.skipUnless(is_sm_100f(), "trtllm-gen attention requires SM100f") @torch.no_grad() @unittest.mock.patch( "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None ) - def test_cuda_graph_decode_high_gqa(self): + def test_cuda_graph_decode_high_gqa(self) -> None: """CUDA graph decode with GQA=8 and real head_dim (E2B-like). Uses E2B real-dims config (hd=256/512, GQA=8) with multi-step @@ -3629,11 +3634,12 @@ def test_cuda_graph_decode_26b_like(self): """26B-like: GQA=2, K=V, hd=256/512.""" self._run_cuda_graph_real_headdim(deepcopy(GEMMA4_26B_REAL_DIMS_CONFIG), "26B") + @unittest.skipUnless(is_sm_100f(), "trtllm-gen attention requires SM100f") @torch.no_grad() @unittest.mock.patch( "tensorrt_llm.runtime.kv_cache_manager_v2._utils.assert_critical", lambda *a, **kw: None ) - def test_cuda_graph_multi_step_trtllm_gen(self): + def test_cuda_graph_multi_step_trtllm_gen(self) -> None: """Multi-step CG decode with trtllm-gen (hd=256/512). Verifies _block_tables update in prepare() works correctly From 6604148cf1e1bd7a287b085419ee375a798ff6ff Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Fri, 21 Aug 2026 08:20:57 -0700 Subject: [PATCH 7/9] [None][fix] Preserve Gemma4 SWA in chunked prefill Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- tensorrt_llm/_torch/models/modeling_gemma4.py | 25 +++++++++++++------ .../_torch/modeling/test_modeling_gemma4.py | 25 +++++++++++++++++++ 2 files changed, 42 insertions(+), 8 deletions(-) diff --git a/tensorrt_llm/_torch/models/modeling_gemma4.py b/tensorrt_llm/_torch/models/modeling_gemma4.py index 9acd5a5e2b3f..f535b17e53bf 100644 --- a/tensorrt_llm/_torch/models/modeling_gemma4.py +++ b/tensorrt_llm/_torch/models/modeling_gemma4.py @@ -1339,11 +1339,10 @@ def get_context_mask( """Build context mask with causal + bidirectional for MM tokens. Returns a [extend_len, prefix_len + extend_len] mask where: - - The first `prefix_len` columns (cached/paged history) are True for - all rows. SWA window enforcement is delegated to the kernel's - window_left clip. Bidirectional MM across the prefix/extend - boundary is NOT supported here; callers must ensure chunk - boundaries do not split a multimodal block. + - The first `prefix_len` columns (cached/paged history) apply the + sliding window using absolute token positions. Bidirectional MM + across the prefix/extend boundary is NOT supported here; callers + must ensure chunk boundaries do not split a multimodal block. - The last `extend_len` columns follow the original causal + (optional) sliding window + MM-bidirectional logic. """ @@ -1360,9 +1359,19 @@ def get_context_mask( causal_mask = causal_mask.masked_fill(token_type_mask, True) if prefix_len > 0: - prefix_block = torch.ones( - extend_len, prefix_len, dtype=causal_mask.dtype, device=device - ) + if ( + effective_sliding_window is not None + and effective_sliding_window < prefix_len + extend_len + ): + query_pos = prefix_len + pos + prefix_pos = torch.arange(prefix_len, device=device) + prefix_block = ( + prefix_pos.unsqueeze(0) > query_pos.unsqueeze(1) - effective_sliding_window + ) + else: + prefix_block = torch.ones( + extend_len, prefix_len, dtype=causal_mask.dtype, device=device + ) causal_mask = torch.cat([prefix_block, causal_mask], dim=1) return causal_mask diff --git a/tests/unittest/_torch/modeling/test_modeling_gemma4.py b/tests/unittest/_torch/modeling/test_modeling_gemma4.py index d464882a7d14..e57baacb94c2 100644 --- a/tests/unittest/_torch/modeling/test_modeling_gemma4.py +++ b/tests/unittest/_torch/modeling/test_modeling_gemma4.py @@ -2337,6 +2337,31 @@ def test_bidirectional_mask_gating(self): # But causal for text tokens self.assertFalse(mask_26b[0, 1].item(), "Text token 0 should NOT attend to 1") + @torch.no_grad() + def test_chunked_context_mask_applies_prefix_window(self) -> None: + """Chunked prefill preserves the full-sequence sliding-window mask.""" + config_dict = deepcopy(GEMMA4_E4B_LIKE_CONFIG) + config_dict["use_bidirectional_attention"] = "vision" + config = Gemma4TextConfig(**config_dict) + model_config = ModelConfig(pretrained_config=config, attn_backend="FLASHINFER") + model = Gemma4ForCausalLM(model_config).to(config.torch_dtype).to("cuda") + + window = 64 + chunk_start = 48 + token_type_ids = torch.zeros(96, dtype=torch.long, device="cuda") + token_type_ids[64:80] = 1 + + full_mask = model.get_context_mask(token_type_ids, effective_sliding_window=window) + chunk_mask = model.get_context_mask( + token_type_ids[chunk_start:], + effective_sliding_window=window, + prefix_len=chunk_start, + ) + + torch.testing.assert_close(chunk_mask, full_mask[chunk_start:]) + self.assertFalse(chunk_mask[-1, 0].item()) + self.assertTrue(chunk_mask[-1, 32].item()) + @torch.no_grad() def test_bidirectional_mask_only_applies_to_sliding_layers(self): """Full-attention layers retain the standard causal mask.""" From 41341c189355e92d96ab280cd7504cae949e407d Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Fri, 21 Aug 2026 21:00:18 -0700 Subject: [PATCH 8/9] [None][docs] Clarify FA2 graph schedule refresh ownership Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- .../_torch/attention_backend/flashinfer.py | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) diff --git a/tensorrt_llm/_torch/attention_backend/flashinfer.py b/tensorrt_llm/_torch/attention_backend/flashinfer.py index 737f8769ad97..648e85eda22a 100644 --- a/tensorrt_llm/_torch/attention_backend/flashinfer.py +++ b/tensorrt_llm/_torch/attention_backend/flashinfer.py @@ -1028,7 +1028,6 @@ def _post_init_with_buffers(self, buffers) -> None: 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._host_paged_kv_last_page_len: Optional[torch.Tensor] = None self._max_num_blocks_per_seq = 0 # VSWA (Variable Sliding Window Attention): models with per-layer @@ -1451,6 +1450,10 @@ def _clean_cached_plans(self, *, defer_plan: bool): elif wrappers.fa2_plan_num_blocks == num_blocks: continue else: + # Graph replay does not re-enter forward_impl. Refresh + # captured FA2 schedules here; each wrapper owns its + # persistent integer plan workspace, while the shared + # float workspace is run scratch. wrappers.is_planned = False self._plan_with_params(plan_params) continue @@ -1696,11 +1699,10 @@ def _to_int32_tensor(arr: np.ndarray) -> torch.Tensor: self._positions[:positions.size(0)].copy_(positions, non_blocking=True) - # Multi-wrapper case (Gemma4 hybrid: different head_dim per layer) - # shares one workspace_buffer; eager plan() would overwrite earlier - # wrappers' workspace, so defer plan() to forward_impl. Single-wrapper - # case (e.g., Llama, Gemma3 uniform head_dim) needs eager plan() here - # because forward_impl cannot plan() during cuda-graph stream capture. + # Defer ordinary multi-wrapper plans to forward_impl. Captured FA2 + # schedule refreshes are handled eagerly by _clean_cached_plans because + # graph replay does not re-enter Python. Single-wrapper models still + # plan eagerly because forward_impl cannot plan during graph capture. active_wrappers = [ pp for pp in self._plan_params_to_wrappers if pp.attention_mask_data is None From f912100363c7e6eba105271833426bdd4ed804ad Mon Sep 17 00:00:00 2001 From: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> Date: Tue, 25 Aug 2026 10:00:45 -0700 Subject: [PATCH 9/9] [None][refactor] Clarify FA2 graph plan refresh lifecycle Move the captured FA2 schedule refresh to one explicit point after prepare finishes updating page metadata. Remove ineffective FA2 _max_kv_len preservation. Signed-off-by: Fanrong Li <23290157+lfr-0531@users.noreply.github.com> --- .../_torch/attention_backend/flashinfer.py | 63 ++++++++++--------- 1 file changed, 32 insertions(+), 31 deletions(-) diff --git a/tensorrt_llm/_torch/attention_backend/flashinfer.py b/tensorrt_llm/_torch/attention_backend/flashinfer.py index 648e85eda22a..c16d2ffaee36 100644 --- a/tensorrt_llm/_torch/attention_backend/flashinfer.py +++ b/tensorrt_llm/_torch/attention_backend/flashinfer.py @@ -1444,25 +1444,30 @@ def _clean_cached_plans(self, *, defer_plan: bool): if plan_params.attention_mask_data is None and plan_params.multi_item_params is None: wrappers = self._plan_params_to_wrappers[plan_params] if wrappers.fa2_plan_num_blocks is not None: - num_blocks = tuple(self.num_blocks[self.num_contexts:]) - if not num_blocks: - wrappers.fa2_plan_num_blocks = None - elif wrappers.fa2_plan_num_blocks == num_blocks: - continue - else: - # Graph replay does not re-enter forward_impl. Refresh - # captured FA2 schedules here; each wrapper owns its - # persistent integer plan workspace, while the shared - # float workspace is run scratch. - wrappers.is_planned = False - self._plan_with_params(plan_params) - continue + continue wrappers.is_planned = False if not defer_plan: self._plan_with_params(plan_params) else: del self._plan_params_to_wrappers[plan_params] + def _refresh_fa2_cuda_graph_plans(self) -> None: + """Refresh captured FA2 schedules after page metadata is finalized.""" + num_blocks = tuple(self.num_blocks[self.num_contexts:]) + for plan_params, wrappers in self._plan_params_to_wrappers.items(): + if (plan_params.attention_mask_data is not None + or plan_params.multi_item_params is not None + or wrappers.fa2_plan_num_blocks is None): + continue + if not num_blocks: + wrappers.fa2_plan_num_blocks = None + elif wrappers.fa2_plan_num_blocks != num_blocks: + # Graph replay does not re-enter forward_impl. Each wrapper + # owns its persistent integer plan workspace, while the shared + # float workspace is run scratch. + wrappers.is_planned = False + self._plan_with_params(plan_params) + def prepare(self) -> None: def _to_int32_tensor(arr: np.ndarray) -> torch.Tensor: @@ -1699,18 +1704,6 @@ def _to_int32_tensor(arr: np.ndarray) -> torch.Tensor: self._positions[:positions.size(0)].copy_(positions, non_blocking=True) - # Defer ordinary multi-wrapper plans to forward_impl. Captured FA2 - # schedule refreshes are handled eagerly by _clean_cached_plans because - # graph replay does not re-enter Python. Single-wrapper models still - # plan eagerly because forward_impl cannot plan during graph capture. - active_wrappers = [ - pp for pp in self._plan_params_to_wrappers - if pp.attention_mask_data is None - ] - defer_plan = len(active_wrappers) > 1 - if not (self._is_separate_kv_draft_view and self.is_cuda_graph): - self._clean_cached_plans(defer_plan=defer_plan) - # Re-plan MLA wrappers outside of forward/capture using the params # cached by prior warmup forwards. Forward still handles first-use or # dtype/shape changes by syncing only on a plan cache miss. @@ -1840,6 +1833,20 @@ def _to_int32_tensor(arr: np.ndarray) -> torch.Tensor: non_blocking=True) if self.num_generations < batch_size: kv_lens_buf[self.num_generations:batch_size].zero_() + + # Refresh captured FA2 schedules only after all page metadata updates. + # Defer ordinary multi-wrapper plans to forward_impl; single-wrapper + # models still plan eagerly because forward_impl cannot plan during + # graph capture. + active_wrappers = [ + pp for pp in self._plan_params_to_wrappers + if pp.attention_mask_data is None + ] + defer_plan = len(active_wrappers) > 1 + if not (self._is_separate_kv_draft_view and self.is_cuda_graph): + self._refresh_fa2_cuda_graph_plans() + self._clean_cached_plans(defer_plan=defer_plan) + if (not self._is_shared_kv_draft_view and not self._is_separate_kv_draft_view and self._draft_metadata is not None): @@ -2072,9 +2079,6 @@ def decode_plan(): if decode_wrapper._backend == 'trtllm-gen': block_tables = self._build_decode_block_tables( plan_params, wrappers) - planned_max_kv_len = (decode_wrapper._max_kv_len - if wrappers.fa2_plan_num_blocks is not None - else None) decode_wrapper.plan( paged_kv_indptr[:self.num_generations + 1], self.paged_kv_indices[self.num_context_blocks:], @@ -2093,9 +2097,6 @@ def decode_plan(): q_len_per_req=plan_params.q_len_per_req, disable_split_kv=False, ) - if planned_max_kv_len is not None: - decode_wrapper._max_kv_len = max(planned_max_kv_len, - decode_wrapper._max_kv_len) if use_graph_tensor_cores and decode_wrapper._backend == 'fa2': wrappers.fa2_plan_num_blocks = tuple( self.num_blocks[self.num_contexts:])