Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 25 additions & 12 deletions tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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"
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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),
)
Comment thread
yuxianq marked this conversation as resolved.
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
Expand Down
83 changes: 78 additions & 5 deletions tests/unittest/_torch/executor/test_per_layer_head_dim.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
import gc
import unittest

import pytest
import torch

import tensorrt_llm
Expand All @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -196,18 +199,18 @@ 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()


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):
Expand Down Expand Up @@ -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()
Loading