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
Original file line number Diff line number Diff line change
Expand Up @@ -310,14 +310,10 @@ def _format_kv_cache_pool_lifecycle_entry(self, layer_id: LayerId, role: DataRol
return super()._format_kv_cache_pool_lifecycle_entry(layer_id, role)

model_layer_idx, attn_type = layer_semantics
attr = self.impl._storage.get_buffer_attr(layer_id, role)
pool_group_id = self.impl._storage.get_pool_group_index(attr.life_cycle_id)
lifecycle = self.impl._life_cycles.get_life_cycle(attr.life_cycle_id)
return (
f"deepseek_role={attn_type.name}, "
f"compress_ratio={self._compress_ratios[model_layer_idx]}, "
f"pool_group_id={int(pool_group_id)}, "
f"lifecycle_id={int(attr.life_cycle_id)}, lifecycle={lifecycle}"
f"{super()._format_kv_cache_pool_lifecycle_entry(layer_id, role)}"
)

def get_buffers(self, layer_idx: int, attn_type: DeepseekV4AttentionType) -> torch.Tensor:
Expand Down Expand Up @@ -934,7 +930,6 @@ def _add_layer(

return KVCacheManagerConfigPy(
tokens_per_block=tokens_per_block,
vocab_size=vocab_size,
cache_tiers=cache_tiers,
max_util_for_resume=kv_cache_config.max_util_for_resume,
swa_scratch_reuse=scratch_reuse_config,
Expand Down
205 changes: 89 additions & 116 deletions tensorrt_llm/_torch/disaggregation/resource/kv_extractor.py
Original file line number Diff line number Diff line change
Expand Up @@ -288,136 +288,109 @@ def _compute_global_layer_ids(manager, lg_idx: int) -> List[int]:
def _build_page_table_v2(manager) -> KVCachePageTable:
"""Build a KVCachePageTable from a KVCacheManagerV2.

Uses the V2 storage layer APIs (pool.slot_address, pool.slot_size,
pool.num_slots) for accurate pool metadata, and stamps each PoolView
with the manager's native role-name strings (``pool_role``) plus the
closed-set ``mapper_kind`` discriminator used by ``build_kv_mapper``.

Important: iterates over life cycles (layer groups), not storage pool
groups. Multiple life cycles with different sliding-window sizes may
share the same underlying storage pool group when their buffer sizes
are identical. The page table must reflect life cycles so that
per-window transfer logic works correctly.
Uses KVCacheManagerV2's public pool_group_descs layout API. A physical
pool group may be shared by several layer groups; layer_groups remains
indexed by layer_group_id while pool_group_idx points at the shared
physical pool group entry.

Each PoolView is stamped with the manager's native role-name strings
(``pool_role``) plus the closed-set ``mapper_kind`` discriminator used
by ``build_kv_mapper``.
"""
from collections import defaultdict

from tensorrt_llm.runtime.kv_cache_manager_v2 import CacheTier

storage = manager.impl._storage
config = manager.impl._init_config

# Find GPU level
gpu_level = 0
for level_idx, cache_tier_config in enumerate(config.cache_tiers):
if cache_tier_config.tier == CacheTier.GPU_MEM:
gpu_level = level_idx
break

# Collect buffer entries keyed by (life_cycle_id, pool_idx).
# Also collect the set of native role-name strings per pool — used as
# ``PoolView.pool_role``, the manager-supplied equivalence label that
# disagg uses to match pools across peers without enumerating roles.
buffer_by_lc_pool: Dict[tuple, list] = defaultdict(list)
native_roles_by_pool: Dict[tuple, set] = defaultdict(set)

for buffer_id, attr in storage._buffer_attr.items():
layer_id, role = buffer_id
lc_id = attr.life_cycle_id
pool_idx = attr.pool_index
pool_key = (int(lc_id), pool_idx)
buffer_by_lc_pool[pool_key].append((layer_id, attr.offset, attr.size))
native_roles_by_pool[pool_key].add(str(role))

# Iterate over life cycles (layer groups), not storage pool groups.
# Multiple layer_groups can share the same storage pool_group when their
# slot_size_list (coalesced buffer sizes) are identical. In that case,
# different layer_groups draw slots from the same physical pool, but a
# slot is exclusively allocated to one layer_group at a time (managed by
# SlotAllocator). Within a slot, each layer_group's buffer offsets start
# from 0 independently — the memory is reused, not concatenated.
# Therefore, slot_bytes / num_layers_for_this_layer_group correctly gives
# the per-layer size, and buffer offsets within a slot are contiguous for
# each layer_group.
num_life_cycles = storage.num_life_cycles
pool_group_storage = storage._levels[gpu_level].storage._pool_groups
config = manager.impl.init_config
pool_group_descs = manager.impl.pool_group_descs

pool_groups: List[PhysicalPoolGroup] = []
storage_pg_to_list_idx: Dict[int, int] = {}
layer_groups: List[LayerGroup] = []

for lc_idx in range(num_life_cycles):
# Resolve the storage pool group for this life cycle.
# storage_pg_idx may be the same for multiple lc_idx values.
storage_pg_idx = storage.get_pool_group_index(lc_idx)
pool_group = pool_group_storage[storage_pg_idx]
num_pools = pool_group.num_pools

# Build PhysicalPoolGroup once per unique storage pool group.
if storage_pg_idx not in storage_pg_to_list_idx:
storage_pg_to_list_idx[storage_pg_idx] = len(pool_groups)
pool_groups.append(
PhysicalPoolGroup(
pools=[
PhysicalPool(
base_address=int(pool_group._pools[pi].slot_address(0)),
slot_bytes=int(pool_group._pools[pi].slot_size),
num_slots=int(pool_group._pools[pi].num_slots),
)
for pi in range(num_pools)
]
)
)
def _window_size_for_layer(internal_layer_id: int):
if internal_layer_id < len(config.layers):
return getattr(config.layers[internal_layer_id], "window_size", None)

# Compute group-level global layer IDs and internal layer IDs
all_internal_layer_ids = list(manager.impl.layer_grouping[lc_idx])
all_global_layer_ids = _compute_global_layer_ids(manager, lc_idx)
if hasattr(manager, "_layer_attn_to_layer_id"):
for (model_layer, _attn_type), layer_id in manager._layer_attn_to_layer_id.items():
if layer_id != internal_layer_id:
continue
local_layer = manager.layer_offsets.get(model_layer)
if local_layer is not None and local_layer < len(config.layers):
return getattr(config.layers[local_layer], "window_size", None)
if model_layer < len(config.layers):
return getattr(config.layers[model_layer], "window_size", None)

local_layers = [
LocalLayer(local_layer_id=int(iid), global_layer_id=int(gid))
for iid, gid in zip(all_internal_layer_ids, all_global_layer_ids)
]
raise ValueError(f"Cannot resolve layer config for internal layer {internal_layer_id}")

pool_views = []
for pool_idx in range(num_pools):
pool_key = (lc_idx, pool_idx)
buffers_info = buffer_by_lc_pool.get(pool_key, [])

# Skip pools that have no buffers for this layer group.
# Multiple life cycles may share the same storage pool group;
# only include pools that actually belong to this life cycle.
if not buffers_info:
continue

pool_views.append(
PoolView(
pool_idx=pool_idx,
buffer_entries=np.array(buffers_info, dtype=BUFFER_ENTRY_DTYPE),
pool_role=frozenset(native_roles_by_pool[pool_key]),
mapper_kind=MapperKind.INDEXED,
)
pool_groups: List[PhysicalPoolGroup] = []
storage_pg_to_list_idx: Dict[int, int] = {}
layer_groups_by_id: List[LayerGroup | None] = [None] * len(manager.impl.layer_grouping)

for pg_desc in pool_group_descs:
storage_pg_idx = int(pg_desc.pool_group_index)
storage_pg_to_list_idx[storage_pg_idx] = len(pool_groups)
pool_groups.append(
PhysicalPoolGroup(
pools=[
PhysicalPool(
base_address=int(pool.base_address),
slot_bytes=int(pool.slot_bytes),
num_slots=int(pg_desc.num_slots),
)
for pool in pg_desc.pools
]
)
)

# Determine layer group metadata.
# For managers with virtual layers, internal layer_ids
# may exceed the length of num_kv_heads_per_layer. Use index 0 as all
# layers within a pool group share the same kv_heads count.
first_local_layer = all_internal_layer_ids[0]
if first_local_layer < len(manager.num_kv_heads_per_layer):
num_kv_heads = manager.num_kv_heads_per_layer[first_local_layer]
else:
num_kv_heads = manager.num_kv_heads_per_layer[0]
life_cycle = manager.impl._life_cycles[lc_idx]
sliding_window_size = life_cycle.window_size
for variant in pg_desc.slot_desc.variants:
layer_group_id = int(variant.layer_group_id)
all_internal_layer_ids = list(manager.impl.layer_grouping[layer_group_id])
all_global_layer_ids = _compute_global_layer_ids(manager, layer_group_id)

local_layers = [
LocalLayer(local_layer_id=int(iid), global_layer_id=int(gid))
for iid, gid in zip(all_internal_layer_ids, all_global_layer_ids)
]

pool_views = []
for pool_idx, coalesced_buffer in enumerate(variant.coalesced_buffers):
entries = []
# Native role-name strings for this pool — used as
# ``PoolView.pool_role``, the manager-supplied equivalence
# label that disagg uses to match pools across peers without
# enumerating roles.
native_roles: set = set()
offset = 0
single_buffer_size = int(coalesced_buffer.single_buffer_size)
for buffer_id in coalesced_buffer.buffer_ids:
entries.append((int(buffer_id.layer_id), offset, single_buffer_size))
native_roles.add(str(buffer_id.role))
offset += single_buffer_size

if entries:
pool_views.append(
PoolView(
pool_idx=pool_idx,
buffer_entries=np.array(entries, dtype=BUFFER_ENTRY_DTYPE),
pool_role=frozenset(native_roles),
mapper_kind=MapperKind.INDEXED,
)
)

layer_groups.append(
AttentionLayerGroup(
first_local_layer = all_internal_layer_ids[0]
if first_local_layer < len(manager.num_kv_heads_per_layer):
num_kv_heads = manager.num_kv_heads_per_layer[first_local_layer]
else:
num_kv_heads = manager.num_kv_heads_per_layer[0]
sliding_window_size = _window_size_for_layer(first_local_layer)

layer_groups_by_id[layer_group_id] = AttentionLayerGroup(
pool_group_idx=storage_pg_to_list_idx[storage_pg_idx],
kv_head_num_per_rank=num_kv_heads,
sliding_window_size=sliding_window_size,
local_layers=local_layers,
pool_views=pool_views,
)
)

layer_groups: List[LayerGroup] = []
for layer_group_id, layer_group in enumerate(layer_groups_by_id):
if layer_group is None:
raise ValueError(f"Missing V2 layer group descriptor for layer group {layer_group_id}")
layer_groups.append(layer_group)

if isinstance(manager, MambaHybridCacheManager):
mamba_layer_group_idx = len(pool_groups)
Expand Down
32 changes: 23 additions & 9 deletions tensorrt_llm/_torch/disaggregation/resource/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,15 +47,29 @@ def get_pool_view_global_layer_ids(
pool_view: PoolView, layer_group: AttentionLayerGroup
) -> List[int]:
"""
Global layer IDs for the layers that appear in *pool_view*, ordered as in
*layer_group.local_layers*.
"""
local_ids_in_pool = get_unique_layers(pool_view)
return [
ll.global_layer_id
for ll in layer_group.local_layers
if ll.local_layer_id in local_ids_in_pool
]
Global layer IDs for the layers that appear in *pool_view*, ordered by their
physical offset within the coalesced buffer (ascending).

The order is derived from the buffer entries' physical offsets rather than
from ``layer_group.local_layers`` order on purpose: the KV transfer maps
layers positionally (a layer's position in this list times the per-layer
slot size gives its byte offset), so the position must reflect the physical
slot layout. Deriving it from offsets keeps the transceiver decoupled from
the KV-cache manager's layer-grouping order (which is an implementation
detail, not an API contract). This mirrors ``get_aggregated_pages``, which
likewise sorts buffers by their offset inside the coalesced buffer.
"""
local_to_global = {ll.local_layer_id: ll.global_layer_id for ll in layer_group.local_layers}
# A layer may contribute several buffer entries (e.g. KEY and VALUE); use the
# smallest offset as that layer's position within the slot.
min_offset: dict[int, int] = {}
for entry in pool_view.buffer_entries:
local_layer_id = int(entry["local_layer_id"])
offset = int(entry["offset"])
if local_layer_id not in min_offset or offset < min_offset[local_layer_id]:
min_offset[local_layer_id] = offset
ordered_local_ids = sorted(min_offset, key=lambda lid: min_offset[lid])
return [local_to_global[lid] for lid in ordered_local_ids]


# -------------------------------------------------------------------------
Expand Down
Loading
Loading