diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py index cc598cbf67ca..98aba2f428d3 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py @@ -1369,15 +1369,6 @@ def create_cold_page_codec(cache_config: object) -> Optional[object]: self._pool_layer_ids_by_role.setdefault( (pool_id, buffer_id.role), buffer_id.layer_id ) - # num_pools is the logical layer-group count. With SWA scratch reuse, - # scratch slot IDs are only valid with per-layer page indices, so the - # attention op sees one virtual pool per local layer while the - # underlying manager can still group layers. - if self.enable_swa_scratch_reuse: - self.num_attention_op_pools = self.num_local_layers - else: - self.num_attention_op_pools = self.num_pools - num_layers = len(config.layers) self.layer_to_pool_mapping_dict: dict[int, int] = { layer_id: self.impl.get_layer_group_id(layer_id) @@ -1492,7 +1483,7 @@ def _build_pool_mapping_tensors(self): kv_cache_pool_pointers_list = [] kv_cache_pool_mapping_list = [] block_scale_pool_pointers_list = [] - if self.enable_swa_scratch_reuse: + if self._use_per_layer_page_tables: for layer_id in typed_range(LayerId(self.num_local_layers)): pool_id = self.impl.get_layer_group_id(layer_id) role_a, _ = self._get_pool_roles(pool_id) @@ -1625,6 +1616,28 @@ def _build_pool_mapping_tensors(self): return kv_cache_pool_pointers, kv_cache_pool_mapping def _prepare_page_table_tensor(self, index_mapper_capacity: int) -> None: + # Keep role-dependent checks in this default page-table initializer: + # custom layouts such as DeepSeek V4 override it with their own roles. + # A lifecycle group can contain different page sizes in separate + # physical pools (e.g. Gemma's heterogeneous head dimensions). Those + # layers cannot share a pool pointer or page-index scale. Use the + # per-layer conversion path, also required for SWA scratch slots, + # while keeping the underlying lifecycle grouping unchanged. + self._use_per_layer_page_tables = self.enable_swa_scratch_reuse or any( + len( + { + self.impl.get_page_stride(layer_id, self._get_pool_roles(pool_id)[0]) + for layer_id in layer_ids + } + ) + > 1 + for pool_id, layer_ids in enumerate(self.impl.layer_grouping) + ) + if self._use_per_layer_page_tables: + self.num_attention_op_pools = self.num_local_layers + else: + self.num_attention_op_pools = self.num_pools + self.kv_cache_pool_pointers, self.kv_cache_pool_mapping = self._build_pool_mapping_tensors() self.index_scales = torch.empty( self.num_pools, dtype=torch.int32, pin_memory=prefer_pinned(), device="cpu" @@ -1658,7 +1671,7 @@ def _prepare_page_table_tensor(self, index_mapper_capacity: int) -> None: pin_memory=prefer_pinned(), device="cpu", ) - if self.enable_swa_scratch_reuse: + if self._use_per_layer_page_tables: self._prepare_swa_scratch_copy_tensors(index_mapper_capacity) def _kv_pool_mapping_offset( @@ -4363,7 +4376,7 @@ def copy_batch_block_offsets( copy_idx = self.index_mapper.get_copy_index(request_ids, num_contexts, beam_width) assert copy_idx.shape[0] == num_seqs - if self.enable_swa_scratch_reuse: + if self._use_per_layer_page_tables: self._copy_batch_block_offsets_per_layer( dst_tensor, request_ids, copy_idx, num_contexts, num_seqs ) diff --git a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_extra_buffers.py b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_extra_buffers.py index d518d288581e..81622c5f0ee6 100644 --- a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_extra_buffers.py +++ b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_extra_buffers.py @@ -158,31 +158,41 @@ def test_default_hook_keeps_nvfp4_scale_buffers(self): del mgr def test_page_table_uses_physical_pool_representative(self): - mgr = KVCacheManagerV2(**_make_kwargs(head_dim=[64, 192, 64, 192])) - real_impl = mgr.impl - try: - pool_id = 0 - physical_layer = mgr._pool_layer_ids_by_role[(pool_id, Role.KEY)] - other_layer = next( - layer_id - for layer_id in real_impl.layer_grouping[pool_id] - if int(layer_id) != int(physical_layer) - ) - impl_proxy = Mock(wraps=real_impl) - impl_proxy.layer_grouping = ((other_layer, physical_layer),) - mgr.impl = impl_proxy + for head_dim in (128, [64, 192, 64, 192]): + with self.subTest(head_dim=head_dim): + mgr = KVCacheManagerV2(**_make_kwargs(head_dim=head_dim)) + real_impl = mgr.impl + try: + pool_id = 0 + physical_layer = mgr._pool_layer_ids_by_role[(pool_id, Role.KEY)] + layers = real_impl.layer_grouping[pool_id] + other_layer = next(layer for layer in layers if layer != physical_layer) + impl_proxy = Mock(wraps=real_impl) + impl_proxy.layer_grouping = ( + (other_layer, *(layer for layer in layers if layer != other_layer)), + ) + mgr.impl = impl_proxy - mgr._prepare_page_table_tensor(index_mapper_capacity=1) + mgr._prepare_page_table_tensor(index_mapper_capacity=1) - self.assertEqual( - impl_proxy.get_mem_pool_base_address.call_args_list[0].args, - (physical_layer, Role.KEY, PageIndexMode.SHARED), - ) - impl_proxy.get_page_index_scale.assert_called_once_with(physical_layer, Role.KEY) - finally: - mgr.impl = real_impl - mgr.shutdown() - del mgr + # Both paths still use a physical representative for group + # metadata. Per-layer pointer lookups can precede it. + shared_key_calls = [ + call.args + for call in impl_proxy.get_mem_pool_base_address.call_args_list + if call.args[1:] == (Role.KEY, PageIndexMode.SHARED) + ] + self.assertEqual( + shared_key_calls[0], + (physical_layer, Role.KEY, PageIndexMode.SHARED), + ) + impl_proxy.get_page_index_scale.assert_called_once_with( + physical_layer, Role.KEY + ) + finally: + mgr.impl = real_impl + mgr.shutdown() + del mgr def test_subclass_registers_index_key_only_on_sparse_layers(self): # Sparse layer convention: layers 0-2 dense (no INDEX_KEY), 3+ sparse diff --git a/tests/unittest/_torch/executor/test_per_layer_head_dim.py b/tests/unittest/_torch/executor/test_per_layer_head_dim.py index 851f56377081..38b1e0a12ef3 100644 --- a/tests/unittest/_torch/executor/test_per_layer_head_dim.py +++ b/tests/unittest/_torch/executor/test_per_layer_head_dim.py @@ -16,6 +16,7 @@ import gc import unittest +import pytest import torch import tensorrt_llm @@ -27,6 +28,7 @@ ) from tensorrt_llm.llmapi.llm_args import KvCacheConfig as KvCacheConfigV2 from tensorrt_llm.mapping import Mapping +from tensorrt_llm.runtime.kv_cache_manager_v2 import PageIndexMode DataType = tensorrt_llm.bindings.DataType CacheType = tensorrt_llm.bindings.internal.batch_manager.CacheType @@ -50,6 +52,7 @@ def _create_kv_cache_manager_v2( kv_cache_config = KvCacheConfigV2( max_tokens=max_tokens, enable_block_reuse=False, + dtype="nvfp4" if dtype == DataType.NVFP4 else "auto", ) return KVCacheManagerV2( kv_cache_config, @@ -196,6 +199,7 @@ def test_per_layer_head_dim_with_equal_buffer_sizes(self): self.assertEqual(bytes_0, 4 * 64 * 2) self.assertEqual(bytes_1, 2 * 128 * 2) self.assertEqual(bytes_0, bytes_1) + self.assertEqual(mgr.num_attention_op_pools, mgr.num_pools) finally: mgr.shutdown() @@ -203,11 +207,10 @@ def test_per_layer_head_dim_with_equal_buffer_sizes(self): class TestPerLayerHeadDimHeterogeneous(unittest.TestCase): """Tests with different buffer sizes per layer. - All tests that create managers with heterogeneous per-layer buffer sizes - are consolidated into a single test method. This is necessary because - CUDA virtual memory addresses from destroyed managers may not be fully - reclaimed, causing subsequent managers to use addresses that lead to - large page offsets (exceeding int32 in pool mapping tensors). + These cases were originally combined to avoid large offsets between + separately allocated pools. Heterogeneous pages now use per-layer base + pointers and page indices, so their mappings no longer depend on the + distance between pools. Explicit GPU cleanup is retained. """ def setUp(self): @@ -336,5 +339,75 @@ def test_per_layer_head_dim_heterogeneous(self): mgr_mixed.shutdown() +@pytest.mark.parametrize("dtype", [DataType.HALF, DataType.FP8, DataType.NVFP4]) +@pytest.mark.parametrize("is_gen", [False, True]) +def test_heterogeneous_page_tables_match_allocated_addresses(dtype, is_gen): + # Gemma's short-sequence configuration puts both head sizes in one + # lifecycle group. Unequal layer counts also give them different page + # index scales, so sharing a group's pointer or scale is incorrect. + mgr = _create_kv_cache_manager_v2(num_layers=6, head_dim=[256] * 5 + [512], dtype=dtype) + try: + request_ids = [11, 22, 33] + token_nums = [1, 9, 17] + requests = mgr.add_dummy_requests(request_ids, token_nums, is_gen=is_gen) + assert requests is not None + request_ids.reverse() + token_nums.reverse() + block_offsets = torch.empty( + mgr.num_attention_op_pools, + len(request_ids), + 2, + mgr.max_blocks_per_seq, + dtype=torch.int32, + device="cuda", + ) + mgr.copy_batch_block_offsets( + block_offsets, + request_ids, + beam_width=1, + num_contexts=0 if is_gen else len(request_ids), + num_seqs=len(request_ids), + ) + torch.cuda.synchronize() + block_offsets = block_offsets.cpu() + for layer_id in range(mgr.num_local_layers): + pool_id, layer_offset = mgr.kv_cache_pool_mapping[layer_id].tolist() + lifecycle_id = mgr.impl.get_layer_group_id(layer_id) + roles = [(Role.KEY, Role.VALUE)] + if dtype == DataType.NVFP4: + roles.append((Role.KEY_BLOCK_SCALE, Role.VALUE_BLOCK_SCALE)) + for buffer_idx, role_pair in enumerate(roles): + pool_pointer = mgr.kv_cache_pool_pointers[pool_id, 0] + if dtype == DataType.NVFP4: + pool_pointer = pool_pointer[buffer_idx] + for kv_idx, role in enumerate(role_pair): + stride = mgr.impl.get_page_stride(layer_id, role) + layer_base = mgr.impl.get_mem_pool_base_address( + layer_id, role, PageIndexMode.SHARED + ) + scale = mgr.impl.get_page_index_scale(layer_id, role) + for seq_idx, (req_id, token_num) in enumerate(zip(request_ids, token_nums)): + base_indices = mgr.kv_cache_map[req_id].get_base_page_indices(lifecycle_id) + num_blocks = (token_num + mgr.tokens_per_block - 1) // mgr.tokens_per_block + page_indices = block_offsets[pool_id, seq_idx, kv_idx, :num_blocks].long() + actual = ( + int(pool_pointer) + + (layer_offset * mgr.kv_factor + page_indices) * stride + ) + allocated_indices = torch.tensor( + list(base_indices[:num_blocks]), dtype=torch.int64 + ) + expected = layer_base + allocated_indices * scale * stride + torch.testing.assert_close( + actual, + expected, + rtol=0, + atol=0, + msg=f"layer={layer_id}, role={role}, request={req_id}", + ) + finally: + mgr.shutdown() + + if __name__ == "__main__": unittest.main()