Skip to content
Draft
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
4 changes: 3 additions & 1 deletion docs/source/features/kv-cache-connector.md
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,7 @@ These methods run on all workers (GPU processes) and interact with the actual GP
* **Description**: Called at initialization **instead of** `register_kv_caches` when the KV cache manager is `KVCacheManagerV2`, whose memory cannot be expressed as one tensor: there is one slot address space per pool and one page-index space per layer group. The default implementation raises, so a connector that does not implement it can only run on V1.
* **Arguments**: `layout` describes the byte ranges that repeat per page slot. Each `KvCacheLayerGroupLayout` carries a tuple of `KvCacheRegion`s, and the bytes for page slot `i` of a region live at `region.base + region.stride * i` for `region.size` bytes, or equivalently at `region.as_tensor()[i]`. Page indices arriving in `RequestData.new_block_ids_by_layer_group` are scoped to a layer group and index that group's regions.
* **Why regions rather than a tensor**: because the ranges are described rather than implied, the same structure covers MLA (a pool simply has no `value` buffer), sliding-window and hybrid models (one layer group per window size), and non-uniform slots such as MiniMax-M3's index-K buffer sitting beside K/V, without any of them being a special case.
* **Replicated roles**: a group's ranges come in two sets. `regions` holds bytes particular to one attention shard, while `replicated_regions` holds bytes identical on every shard. MiniMax-M3's index-K is computed from a replicated projection, so all TP ranks hold the same values. The manager declares which roles those are through `get_replicated_roles()`, itself derived from the `get_disagg_role_mapper_kinds()` declaration that the native disaggregation path already uses. V2 may interleave the two classes within one pool, so the split is produced by aggregating each class separately rather than by slicing a merged range; a connector that ignores `replicated_regions` will simply not transfer those bytes.

* **`start_load_kv(self, stream: torch.cuda.Stream)`**
* **Description**: Initiates the loading of KV blocks from the external source into the GPU memory.
Expand Down Expand Up @@ -143,7 +144,8 @@ explicitly opt in.
The built-in `mooncake-store` adapter shares one unsharded attention namespace
across ADP owners. Each owner opens its own store client and contributes its
configured segment to the common master. TP uses separate keys for each
attention shard and still requires all shards for a prefix hit. A disaggregated
attention shard and still requires all shards for a prefix hit, except for roles
the manager declares replicated, which the whole TP group stores once. A disaggregated
DEP4 prefill / TEP8 decode deployment attaches the store connector to prefill;
decode can donate host memory while keeping its native KV transceiver. The
store does not convert attention shard layouts during prefill-to-decode handoff.
Expand Down
91 changes: 68 additions & 23 deletions tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_layout.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,13 @@
Because ranges are described rather than implied, this covers MLA (a pool simply
has no VALUE buffer), sliding-window attention and hybrid models (one layer group
per window size) without any of them being special cases.

A group's ranges are split into two sets. Most roles hold bytes particular to one
attention shard, but a side cache may hold bytes identical on every shard, such
as MiniMax-M3's index-K. The manager declares which roles those are, and they
are described as replicated_regions so a connector can store them once for a TP
group instead of once per rank. V2 may interleave the two classes within a pool,
so each class is aggregated separately rather than sliced out of a merged range.
"""

from dataclasses import dataclass
Expand Down Expand Up @@ -120,13 +127,22 @@ class KvCacheLayerGroupLayout:
layer_ids: Tuple[int, ...]
#: Attention window for this group, or None for full attention.
window_size: Optional[int]
#: Regions holding bytes specific to this attention shard.
regions: Tuple[KvCacheRegion, ...]
#: Regions holding bytes identical on every attention shard, as named by
#: the manager's get_replicated_roles(). Empty unless a role is declared.
replicated_regions: Tuple[KvCacheRegion, ...] = ()

@property
def bytes_per_page(self) -> int:
"""Total bytes this group occupies for a single page slot."""
"""Bytes this group's shard-specific regions occupy for one page slot."""
return sum(region.size for region in self.regions)

@property
def replicated_bytes_per_page(self) -> int:
"""Bytes this group's replicated regions occupy for one page slot."""
return sum(region.size for region in self.replicated_regions)


@dataclass(frozen=True)
class KvCacheLayout:
Expand Down Expand Up @@ -187,13 +203,53 @@ def _window_size(init_config, local_layer_id: int) -> Optional[int]:
return None if window is None else int(window)


def _regions_for(
impl,
buffer_ids: List,
layer_group_id: int,
num_slots: int,
global_by_local: Dict[int, int],
) -> Tuple[KvCacheRegion, ...]:
"""Aggregate buffer_ids into the regions they occupy in one slot.

get_aggregated_pages coalesces only buffers adjacent *within the set it is
given*, so passing a subset yields regions covering exactly that subset.
That is what keeps the two region classes separable.
"""
regions: List[KvCacheRegion] = []
for desc in impl.get_aggregated_pages(buffer_ids):
if int(desc.layer_group_id) != layer_group_id:
continue
regions.append(
KvCacheRegion(
base=int(desc.base),
size=int(desc.size),
stride=int(desc.stride),
num_slots=num_slots,
buffers=tuple(
KvCacheBufferRef(
layer_id=global_by_local[int(b.id.layer_id)],
role=str(b.id.role),
expansion=int(b.expansion),
)
for b in desc.buffers
),
)
)
return tuple(regions)


def build_kv_cache_layout_v2(manager: "KVCacheManagerV2") -> KvCacheLayout:
"""Describe a ``KVCacheManagerV2``'s GPU pools for a KV connector.

Built only from V2's public layout API -- ``layer_grouping``,
``all_buffer_ids``, ``get_aggregated_pages`` and ``pool_group_descs``. No
private storage state is touched, and no assumption is made about dimension
order, kv factor, or the number of pools.

Roles the manager declares replicated are described as a separate region
set. A manager that declares none leaves replicated_regions empty and every
buffer in regions.
"""
impl = manager.impl
init_config = impl.init_config
Expand All @@ -210,6 +266,11 @@ def build_kv_cache_layout_v2(manager: "KVCacheManagerV2") -> KvCacheLayout:
for buffer_id in impl.all_buffer_ids:
buffers_by_layer.setdefault(int(buffer_id.layer_id), []).append(buffer_id)

# Compare role names rather than DataRole values: a region records the
# manager's native role string, which keeps this free of the
# disaggregation types used to express the declaration.
replicated_roles = {str(role) for role in manager.get_replicated_roles()}

groups: List[KvCacheLayerGroupLayout] = []
for layer_group_id, local_layer_ids in enumerate(impl.layer_grouping):
local_layer_ids = [int(lid) for lid in local_layer_ids]
Expand All @@ -220,34 +281,18 @@ def build_kv_cache_layout_v2(manager: "KVCacheManagerV2") -> KvCacheLayout:
global_by_local = dict(zip(local_layer_ids, _global_layer_ids(manager, local_layer_ids)))

buffer_ids = [b for lid in local_layer_ids for b in buffers_by_layer.get(lid, ())]

regions: List[KvCacheRegion] = []
for desc in impl.get_aggregated_pages(buffer_ids):
if int(desc.layer_group_id) != layer_group_id:
continue
regions.append(
KvCacheRegion(
base=int(desc.base),
size=int(desc.size),
stride=int(desc.stride),
num_slots=num_slots,
buffers=tuple(
KvCacheBufferRef(
layer_id=global_by_local[int(b.id.layer_id)],
role=str(b.id.role),
expansion=int(b.expansion),
)
for b in desc.buffers
),
)
)
sharded_ids = [b for b in buffer_ids if str(b.role) not in replicated_roles]
replicated_ids = [b for b in buffer_ids if str(b.role) in replicated_roles]

groups.append(
KvCacheLayerGroupLayout(
layer_group_id=layer_group_id,
layer_ids=tuple(global_by_local[lid] for lid in local_layer_ids),
window_size=_window_size(init_config, local_layer_ids[0]),
regions=tuple(regions),
regions=_regions_for(impl, sharded_ids, layer_group_id, num_slots, global_by_local),
replicated_regions=_regions_for(
impl, replicated_ids, layer_group_id, num_slots, global_by_local
),
)
)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,10 @@
and parallel layout because `build_kv_cache_layout_v2` derives it from the
allocator's own aggregation. `bytes_per_page` goes into the key namespace so a
geometry change cannot be read as a valid page.

A layer group yields two pages rather than one, since its shard-specific and
replicated regions are keyed separately. Groups with no replicated role report
an empty second page, which the worker skips.
"""

from typing import Dict, Iterable, List, Optional, Sequence, Tuple
Expand Down Expand Up @@ -91,8 +95,8 @@ def merge_intervals(intervals: Iterable[Tuple[int, int]]) -> List[Tuple[int, int

A range may not be registered twice, but several regions routinely live
inside one pool allocation: sliding-window layer groups share it, and a
non-uniform slot such as MiniMax-M3's index-K beside K/V splits one pool
into several regions. Merging first spares the caller that distinction.
pool holding both region classes (MiniMax-M3's index-K beside K/V) splits
into interleaved regions. Merging first spares the caller that distinction.
"""
ordered = sorted((int(start), int(end)) for start, end in intervals if end > start)
merged: List[Tuple[int, int]] = []
Expand All @@ -111,8 +115,11 @@ class PageAddressing:
def __init__(self, layout: KvCacheLayout):
self._layout = layout
self._regions: Dict[int, Tuple[KvCacheRegion, ...]] = {}
self._replicated_regions: Dict[int, Tuple[KvCacheRegion, ...]] = {}
self._origins: Dict[int, Tuple[int, ...]] = {}
self._replicated_origins: Dict[int, Tuple[int, ...]] = {}
self._bytes_per_page: Dict[int, int] = {}
self._replicated_bytes_per_page: Dict[int, int] = {}
self._num_slots: Dict[int, int] = {}
boundary = layout.gpu_pool_mapping_bytes
for group in layout.groups:
Expand All @@ -122,15 +129,25 @@ def __init__(self, layout: KvCacheLayout):
"is nothing for the connector to transfer"
)
self._regions[group.layer_group_id] = group.regions
self._replicated_regions[group.layer_group_id] = group.replicated_regions
# Asked once per region here rather than per transfer: a region's
# reservation is fixed for the lifetime of the pools.
# reservation is fixed for the lifetime of the pools. Each class
# carries its own origins because the two interleave inside a pool,
# and a region has to be cut against the reservation it sits in.
self._origins[group.layer_group_id] = tuple(
mapping_origin(region.base, boundary) for region in group.regions
)
self._replicated_origins[group.layer_group_id] = tuple(
mapping_origin(region.base, boundary) for region in group.replicated_regions
)
self._bytes_per_page[group.layer_group_id] = group.bytes_per_page
self._replicated_bytes_per_page[group.layer_group_id] = group.replicated_bytes_per_page
# Regions of a group come from the same pool group and so share a
# slot count. Disagreement would make the page index ambiguous.
slot_counts = {region.num_slots for region in group.regions}
# Replicated regions share that space, so they are held to it too.
slot_counts = {
region.num_slots for region in (*group.regions, *group.replicated_regions)
}
if len(slot_counts) != 1:
raise ValueError(
f"layer group {group.layer_group_id} mixes slot counts "
Expand All @@ -156,7 +173,12 @@ def mapping_origins(self) -> Tuple[int, ...]:
from the driver and the boundaries are being assumed; see
`mapping_origin`.
"""
origins = {origin for group in self._origins.values() for origin in group}
origins = {
origin
for by_group in (self._origins, self._replicated_origins)
for group in by_group.values()
for origin in group
}
return tuple(sorted(origins))

@property
Expand All @@ -170,15 +192,23 @@ def tokens_per_block(self) -> int:
return self._layout.tokens_per_block

def bytes_per_page(self, layer_group_id: int) -> int:
"""Total payload size of one page of `layer_group_id`."""
"""Payload size of one shard-specific page of `layer_group_id`."""
return self._bytes_per_page[layer_group_id]

def replicated_bytes_per_page(self, layer_group_id: int) -> int:
"""Payload size of one replicated page of `layer_group_id`, or 0."""
return self._replicated_bytes_per_page[layer_group_id]

def has_replicated(self, layer_group_id: int) -> bool:
"""Whether `layer_group_id` holds any replicated-role bytes."""
return bool(self._replicated_regions[layer_group_id])

def num_slots(self, layer_group_id: int) -> int:
"""Number of page slots addressable in `layer_group_id`."""
return self._num_slots[layer_group_id]

def buffers(self, layer_group_id: int, page_index: int) -> Tuple[List[int], List[int]]:
"""Addresses and sizes of one page, in the order they concatenate.
"""Addresses and sizes of one shard-specific page, in concatenation order.

A region whose bytes for this page cross a GPU pool mapping boundary
contributes several consecutive buffers rather than one, which leaves
Expand All @@ -191,7 +221,44 @@ def buffers(self, layer_group_id: int, page_index: int) -> Tuple[List[int], List
Returns:
Parallel lists of device addresses and byte counts.
"""
regions = self._regions[layer_group_id]
return self._locate(
self._regions[layer_group_id],
self._origins[layer_group_id],
layer_group_id,
page_index,
)

def replicated_buffers(
self, layer_group_id: int, page_index: int
) -> Tuple[List[int], List[int]]:
"""Addresses and sizes of one replicated page, in concatenation order.

The bytes match on every rank but the addresses do not, since each rank
names the copy in its own memory. A region crossing a GPU pool mapping
boundary is split the same way as in `buffers`.

Args:
layer_group_id: Layer group the page index is scoped to.
page_index: Page slot index within that group.

Returns:
Parallel lists of device addresses and byte counts. Both are empty
when the group holds no replicated roles.
"""
return self._locate(
self._replicated_regions[layer_group_id],
self._replicated_origins[layer_group_id],
layer_group_id,
page_index,
)

def _locate(
self,
regions: Sequence[KvCacheRegion],
origins: Sequence[int],
layer_group_id: int,
page_index: int,
) -> Tuple[List[int], List[int]]:
num_slots = self._num_slots[layer_group_id]
if not 0 <= page_index < num_slots:
raise IndexError(
Expand All @@ -201,7 +268,7 @@ def buffers(self, layer_group_id: int, page_index: int) -> Tuple[List[int], List
boundary = self._layout.gpu_pool_mapping_bytes
addresses: List[int] = []
sizes: List[int] = []
for region, origin in zip(regions, self._origins[layer_group_id]):
for region, origin in zip(regions, origins):
start = region.base + region.stride * page_index
for address, size in split_at_boundaries(start, region.size, boundary, origin):
addresses.append(address)
Expand All @@ -214,7 +281,9 @@ def registration_ranges(self) -> List[Tuple[int, int]]:
A region's slots are strided rather than packed, so its range spans
from the first slot to the end of the last. Registering the whole span
is what makes every slot's address valid for RDMA, and merging keeps a
shared pool from being registered once per region.
shared pool from being registered once per region. Both region classes
are covered, since replicated bytes are transferred like any other and
an unregistered range cannot be transferred at all.

Merged spans are then cut at GPU pool mapping boundaries, measured from
the reservation each pool was mapped into, since a registration covering
Expand All @@ -227,7 +296,14 @@ def registration_ranges(self) -> List[Tuple[int, int]]:
# group is the same cover as merging all the spans together.
spans_by_origin: Dict[int, List[Tuple[int, int]]] = {}
for layer_group_id, regions in self._regions.items():
for region, origin in zip(regions, self._origins[layer_group_id]):
paired = (
*zip(regions, self._origins[layer_group_id]),
*zip(
self._replicated_regions[layer_group_id],
self._replicated_origins[layer_group_id],
),
)
for region, origin in paired:
span_end = region.base + region.stride * (region.num_slots - 1) + region.size
spans_by_origin.setdefault(origin, []).append((region.base, span_end))

Expand All @@ -248,6 +324,8 @@ def describe(self) -> str:
f"layers={len(group.layer_ids)}, "
f"regions={len(group.regions)}, "
f"bytes/page={group.bytes_per_page}, "
f"replicated_regions={len(group.replicated_regions)}, "
f"replicated_bytes/page={group.replicated_bytes_per_page}, "
f"slots={self._num_slots[group.layer_group_id]}, "
f"window={group.window_size})"
for group in self._layout.groups
Expand Down
Loading
Loading