diff --git a/docs/source/features/kv-cache-connector.md b/docs/source/features/kv-cache-connector.md index 7425e4b2db1e..bf7e24448365 100644 --- a/docs/source/features/kv-cache-connector.md +++ b/docs/source/features/kv-cache-connector.md @@ -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. @@ -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. diff --git a/tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_layout.py b/tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_layout.py index 824f629daddf..c27af2e24ff0 100644 --- a/tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_layout.py +++ b/tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_layout.py @@ -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 @@ -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: @@ -187,6 +203,42 @@ 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. @@ -194,6 +246,10 @@ def build_kv_cache_layout_v2(manager: "KVCacheManagerV2") -> KvCacheLayout: ``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 @@ -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] @@ -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 + ), ) ) diff --git a/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/addressing.py b/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/addressing.py index 3cabc6b3163c..c4615aa91105 100644 --- a/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/addressing.py +++ b/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/addressing.py @@ -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 @@ -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]] = [] @@ -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: @@ -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 " @@ -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 @@ -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 @@ -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( @@ -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) @@ -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 @@ -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)) @@ -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 diff --git a/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/keys.py b/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/keys.py index 036f03af63ba..34e7788e137c 100644 --- a/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/keys.py +++ b/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/keys.py @@ -23,6 +23,10 @@ that decides what the stored bytes mean: the model, the shard that produced them, the layer group inside that shard, the tokens each page holds and how many bytes a page is. A change to any of those reads as a miss rather than as garbage. + +The shard component has two forms. A page whose bytes depend on the shard is +named per rank, while one whose bytes are identical on every rank is named +`REPLICATED_SHARD_KEY` instead, giving a TP group a single shared key. """ import hashlib @@ -34,6 +38,8 @@ "KeyNamespace", "ReuseScope", "HASH_DIGEST_BYTES", + "REPLICATED_SHARD_KEY", + "sharded_shard_key", ] #: 128 bits. Collisions decide whether one request reads another's KV, so the @@ -139,17 +145,31 @@ def extend(self, tokens: Sequence[int]) -> Sequence[bytes]: return self._hashes +#: Shard component for pages whose bytes do not depend on the attention shard. +#: A literal cannot collide with a sharded component, which is always +#: `wr`. +REPLICATED_SHARD_KEY = "replicated" + + +def sharded_shard_key(rank: int, world_size: int) -> str: + """The shard component naming one attention shard. + + Under ADP all owners pass rank 0 of 1, since each holds complete attention + KV. TP keeps distinct shard keys: rank 3 of 8 can hold different heads than + rank 3 of 4. + """ + return f"w{world_size}r{rank}" + + @dataclass(frozen=True) class KeyNamespace: """The part of a store key that is fixed for one shard and layer group.""" namespace: str model_key: str - #: Attention shard rank and count. Under ADP all owners use rank 0 of 1, - #: since each holds complete attention KV. TP keeps distinct shard keys: - #: rank 3 of 8 can hold different heads than rank 3 of 4. - rank: int - world_size: int + #: Which shard's bytes these are, from `sharded_shard_key` for a page whose + #: content depends on its producer, `REPLICATED_SHARD_KEY` otherwise. + shard_key: str layer_group_id: int tokens_per_block: int bytes_per_page: int @@ -159,7 +179,7 @@ def prefix(self) -> str: """The literal string every key in this namespace starts with.""" return ( f"{self.namespace}/{self.model_key}" - f"/w{self.world_size}r{self.rank}" + f"/{self.shard_key}" f"/lg{self.layer_group_id}" f"/t{self.tokens_per_block}b{self.bytes_per_page}" ) diff --git a/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/worker.py b/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/worker.py index f38c368c7ade..3d05df28ebab 100644 --- a/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/worker.py +++ b/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/worker.py @@ -32,6 +32,10 @@ A capacity-only worker opens its handle and stops there, with no layout, no buffer registration and no save thread, so a node can lend host memory to the pool without an HCA that can pin GPU pages. + +A layer group moves as two pages. The shard-specific one is keyed per rank; the +replicated one is keyed once for the whole TP group. One rank writes that shared +copy and every rank reads it, since each needs the bytes in its own GPU memory. """ import os @@ -52,7 +56,7 @@ from .addressing import PageAddressing from .config import CONFIG_PATH_ENV, MooncakeStoreConnectorConfig, pool_config from .gpudirect import REGISTRATION_DEBUG_ENV, format_diagnosis -from .keys import KeyNamespace +from .keys import REPLICATED_SHARD_KEY, KeyNamespace, sharded_shard_key from .ledger import record_segment from .metadata import MooncakeStoreMetadata, RequestTransfers from .staging import ( @@ -193,6 +197,9 @@ def __init__(self, llm_args: TorchLlmArgs) -> None: self._addressing: Optional[PageAddressing] = None #: Namespaces for this attention shard, shared by compatible ADP owners. self._namespaces: Dict[int, KeyNamespace] = {} + # Rank-independent namespaces, present only for the layer groups that + # hold a replicated role. + self._replicated_namespaces: Dict[int, KeyNamespace] = {} # A TP hit requires every attention shard. ADP has one complete shard, # so neither content identity nor lookup depends on unrelated owners. self._peer_namespaces: Dict[int, Tuple[KeyNamespace, ...]] = {} @@ -343,12 +350,24 @@ def register_kv_cache_layout(self, layout: KvCacheLayout) -> None: for layer_group_id in addressing.layer_group_ids: bytes_per_page = addressing.bytes_per_page(layer_group_id) self._namespaces[layer_group_id] = self._namespace( - self._attention_rank, layer_group_id, bytes_per_page + sharded_shard_key(self._attention_rank, self._attention_world_size), + layer_group_id, + bytes_per_page, ) self._peer_namespaces[layer_group_id] = tuple( - self._namespace(rank, layer_group_id, bytes_per_page) + self._namespace( + sharded_shard_key(rank, self._attention_world_size), + layer_group_id, + bytes_per_page, + ) for rank in range(self._attention_world_size) ) + if addressing.has_replicated(layer_group_id): + self._replicated_namespaces[layer_group_id] = self._namespace( + REPLICATED_SHARD_KEY, + layer_group_id, + addressing.replicated_bytes_per_page(layer_group_id), + ) if self._config.role.saves: self._save_thread = threading.Thread( @@ -379,8 +398,13 @@ def _open_staging(self, addressing: PageAddressing) -> None: Only the directions this role drives get a pool, since each one costs a pinned allocation of its own. """ + # Both page classes pass through the same slots, so a slot has to hold + # the larger of the two. max_bytes_per_page = max( - addressing.bytes_per_page(layer_group_id) + max( + addressing.bytes_per_page(layer_group_id), + addressing.replicated_bytes_per_page(layer_group_id), + ) for layer_group_id in addressing.layer_group_ids ) slot_bytes, num_slots = plan_slot_geometry( @@ -412,12 +436,11 @@ def _open_staging(self, addressing: PageAddressing) -> None: f"more. Lower transfer_batch_size to make the reduction explicit." ) - def _namespace(self, rank: int, layer_group_id: int, bytes_per_page: int) -> KeyNamespace: + def _namespace(self, shard_key: str, layer_group_id: int, bytes_per_page: int) -> KeyNamespace: return KeyNamespace( namespace=self._config.namespace, model_key=self._model_key, - rank=rank, - world_size=self._attention_world_size, + shard_key=shard_key, layer_group_id=layer_group_id, tokens_per_block=self._addressing.tokens_per_block, bytes_per_page=bytes_per_page, @@ -435,13 +458,28 @@ def is_registered(self) -> bool: """Whether a KV cache layout has been registered yet.""" return self._addressing is not None + @property + def _owns_replicated_saves(self) -> bool: + """Whether this rank writes the replicated pages of its group. + + Naming a single owner keeps ranks holding identical bytes from racing + to write the same key. + + Under ADP every owner reports attention rank 0, which is deliberate: + owners there serve different requests and so mostly produce different + keys, and the existence filter in `_put` collapses the overlap. Gating + them on rank would drop the pages of every owner but one. + """ + return self._attention_rank == 0 + def count_prefix_hit(self, block_hashes: Sequence[bytes]) -> int: """How many leading blocks of `block_hashes` are fully present. A block counts only when every layer group and attention shard has its - page. ADP owners share the one unsharded representation; TP requires - all shards, a prefix being replayed as a whole. The scan stops at the - first incomplete block, a later hit being unusable on its own. + page, plus the one replicated page a layer group may carry. ADP owners + share the one unsharded representation; TP requires all shards, a + prefix being replayed as a whole. The scan stops at the first + incomplete block, a later hit being unusable on its own. Args: block_hashes: Candidate hashes in block order. @@ -456,6 +494,11 @@ def count_prefix_hit(self, block_hashes: Sequence[bytes]) -> int: for block_hash in block_hashes: for namespaces in self._peer_namespaces.values(): keys.extend(namespace.key(block_hash) for namespace in namespaces) + # One key for the whole group, so a block whose replicated page was + # dropped reads as a miss on every rank. + keys.extend( + namespace.key(block_hash) for namespace in self._replicated_namespaces.values() + ) keys_per_block = len(keys) // len(block_hashes) try: @@ -494,6 +537,15 @@ def start_load_kv(self, stream: torch.cuda.Stream): self._reraise_save_error() keys, addresses, sizes, total_pages = self._resolve(metadata.loads) + # Appending rather than loading separately keeps both page classes + # inside the same transfer batches. + rep_keys, rep_addresses, rep_sizes, rep_pages = self._resolve( + metadata.loads, replicated=True + ) + keys += rep_keys + addresses += rep_addresses + sizes += rep_sizes + total_pages += rep_pages if not keys: return @@ -656,6 +708,11 @@ def _drain_saves(self) -> None: def _put(self, transfers: Sequence[RequestTransfers]) -> None: keys, addresses, sizes, _ = self._resolve(transfers) + if self._owns_replicated_saves: + rep_keys, rep_addresses, rep_sizes, _ = self._resolve(transfers, replicated=True) + keys += rep_keys + addresses += rep_addresses + sizes += rep_sizes if not keys: return @@ -706,9 +763,20 @@ def _put(self, transfers: Sequence[RequestTransfers]) -> None: # ---- shared ---- def _resolve( - self, transfers: Sequence[RequestTransfers] + self, transfers: Sequence[RequestTransfers], replicated: bool = False ) -> Tuple[List[str], List[List[int]], List[List[int]], int]: - """Expand per-request page transfers into parallel store call arguments.""" + """Expand per-request page transfers into parallel store call arguments. + + Args: + transfers: Pages to move, as the scheduler reported them. + replicated: Resolve the replicated page of each layer group rather + than the shard-specific one. Groups holding no replicated role + contribute nothing. + + Returns: + Parallel lists of keys, per-key buffer addresses and per-key buffer + sizes, plus the page count. + """ if self._addressing is None: raise RuntimeError("KV cache layout has not been registered") keys: List[str] = [] @@ -717,15 +785,23 @@ def _resolve( pages = 0 for entry in transfers: for page in entry.pages: - namespace = self._namespaces.get(page.layer_group_id) - if namespace is None: - raise KeyError( - f"layer group {page.layer_group_id} is not in the registered " - "layout; the scheduler and worker disagree about the model" + if replicated: + namespace = self._replicated_namespaces.get(page.layer_group_id) + if namespace is None: + continue + page_addresses, page_sizes = self._addressing.replicated_buffers( + page.layer_group_id, page.page_index + ) + else: + namespace = self._namespaces.get(page.layer_group_id) + if namespace is None: + raise KeyError( + f"layer group {page.layer_group_id} is not in the registered " + "layout; the scheduler and worker disagree about the model" + ) + page_addresses, page_sizes = self._addressing.buffers( + page.layer_group_id, page.page_index ) - page_addresses, page_sizes = self._addressing.buffers( - page.layer_group_id, page.page_index - ) keys.append(namespace.key(page.block_hash)) addresses.append(page_addresses) sizes.append(page_sizes) diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py index 1d898f454028..9e925e82c175 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py @@ -1921,6 +1921,24 @@ def get_disagg_role_mapper_kinds(self) -> dict[DataRole, MapperKind]: """ return {Role.ALL: MapperKind.INDEXED, Role.INDEX_KEY: MapperKind.REPLICATED} + def get_replicated_roles(self) -> frozenset[DataRole]: + """Roles whose bytes are identical on every attention shard. + + Derived from :meth:`get_disagg_role_mapper_kinds` so a manager declares + replication once and every consumer agrees. Role.ALL is the fallback + for roles without an entry rather than a role of its own, so it is + never reported here. + + A KV connector uses this to give replicated bytes a single + rank-independent store key. A role with no registered buffers + contributes nothing, so declaring one is harmless. + """ + return frozenset( + role + for role, kind in self.get_disagg_role_mapper_kinds().items() + if role != Role.ALL and kind is MapperKind.REPLICATED + ) + @property def blocks_in_primary_pool(self) -> int: """ diff --git a/tests/unittest/_torch/executor/test_kv_cache_layout.py b/tests/unittest/_torch/executor/test_kv_cache_layout.py index 83a7d8116ea9..a6c8ecf20dac 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_layout.py +++ b/tests/unittest/_torch/executor/test_kv_cache_layout.py @@ -25,9 +25,10 @@ KvCacheRegion, build_kv_cache_layout_v2, ) -from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2 +from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2, Role 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 BufferConfig DataType = tensorrt_llm.bindings.DataType CacheType = tensorrt_llm.bindings.internal.batch_manager.CacheType @@ -103,6 +104,22 @@ def test_bytes_per_page_sums_regions(self): regions=(self._region(size=256), self._region(base=8192, size=128)), ) self.assertEqual(group.bytes_per_page, 384) + # A group without a replicated role carries no second page. + self.assertEqual(group.replicated_regions, ()) + self.assertEqual(group.replicated_bytes_per_page, 0) + + def test_the_two_page_classes_are_sized_independently(self): + group = KvCacheLayerGroupLayout( + layer_group_id=0, + layer_ids=(0, 1), + window_size=None, + regions=(self._region(size=256), self._region(base=8192, size=128)), + replicated_regions=(self._region(base=16384, size=64),), + ) + # The replicated bytes are a page of their own, since a connector keys + # and transfers them separately. + self.assertEqual(group.bytes_per_page, 384) + self.assertEqual(group.replicated_bytes_per_page, 64) def test_layout_lookup_by_group_and_layer(self): group_a = KvCacheLayerGroupLayout(0, (0, 2), None, ()) @@ -272,6 +289,65 @@ def test_uniform_model_yields_one_full_slot_region(self): [ref.role for ref in region.buffers], ["key", "value"] * 4, ) + # The base manager declares INDEX_KEY replicated but registers no + # such buffer, so every region stays shard-specific. + self.assertEqual(group.replicated_regions, ()) + self.assertEqual(group.replicated_bytes_per_page, 0) + finally: + mgr.shutdown() + del mgr + + def test_a_replicated_role_is_described_as_its_own_regions(self): + # A manager registering an index-K buffer whose per-block size equals + # K/V's gets it coalesced into the same pool, interleaved per layer as + # K, V, INDEX_KEY. The two classes must still come out separated. + num_layers = 4 + kwargs = _make_kwargs(num_layers=num_layers) + bytes_per_block = ( + kwargs["num_kv_heads"] + * kwargs["head_dim"] + * kwargs["tokens_per_block"] + * torch.tensor([], dtype=torch.float16).element_size() + ) + + class _IndexKeyManager(KVCacheManagerV2): + def _extra_buffers_per_layer(self, *, tokens_per_block): + return { + layer: [BufferConfig(role=Role.INDEX_KEY, size=bytes_per_block)] + for layer in range(num_layers) + } + + mgr = _IndexKeyManager(**kwargs) + try: + layout = build_kv_cache_layout_v2(mgr) + self.assertEqual(len(layout.groups), 1) + group = layout.groups[0] + + sharded_roles = [ref.role for region in group.regions for ref in region.buffers] + replicated_roles = [ + ref.role for region in group.replicated_regions for ref in region.buffers + ] + self.assertEqual(sharded_roles, ["key", "value"] * num_layers) + self.assertEqual(replicated_roles, ["index_key"] * num_layers) + + # Interleaving means neither class is one contiguous run, so each + # layer contributes its own region. + self.assertEqual(len(group.regions), num_layers) + self.assertEqual(len(group.replicated_regions), num_layers) + + # Together the two pages still account for the whole slot, and + # neither overlaps the other. + pool = list(mgr.impl.pool_group_descs)[0].pools[0] + self.assertEqual( + group.bytes_per_page + group.replicated_bytes_per_page, + int(pool.slot_bytes), + ) + spans = sorted( + (region.base, region.base + region.size) + for region in (*group.regions, *group.replicated_regions) + ) + for (_, prev_end), (next_start, _) in zip(spans, spans[1:]): + self.assertLessEqual(prev_end, next_start) finally: mgr.shutdown() del mgr diff --git a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py index e45be30dcab0..7b760f174da6 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py +++ b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py @@ -557,3 +557,38 @@ def test_disagg_role_mapper_kinds_default_to_indexed(): Role.ALL: MapperKind.INDEXED, Role.INDEX_KEY: MapperKind.REPLICATED, } + + +def test_replicated_roles_follow_the_disagg_declaration(): + from tensorrt_llm._torch.disaggregation.resource.page import MapperKind + from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import Role + + manager = object.__new__(KVCacheManagerV2) + assert manager.get_replicated_roles() == frozenset({Role.INDEX_KEY}) + + # A subclass declaring a new replicated side cache is picked up on its own. + extra = Role.KEY_BLOCK_SCALE + manager.get_disagg_role_mapper_kinds = lambda: { + Role.ALL: MapperKind.INDEXED, + Role.INDEX_KEY: MapperKind.REPLICATED, + extra: MapperKind.REPLICATED, + } + assert manager.get_replicated_roles() == frozenset({Role.INDEX_KEY, extra}) + + # Sharded roles are excluded, and a manager declaring none reports none. + manager.get_disagg_role_mapper_kinds = lambda: { + Role.ALL: MapperKind.INDEXED, + Role.INDEX_KEY: MapperKind.NHD, + } + assert manager.get_replicated_roles() == frozenset() + + +def test_replicated_roles_never_include_the_fallback(): + from tensorrt_llm._torch.disaggregation.resource.page import MapperKind + from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import Role + + # Role.ALL names the fallback for roles without an entry, not a role whose + # buffers exist, so it must never be reported replicated. + manager = object.__new__(KVCacheManagerV2) + manager.get_disagg_role_mapper_kinds = lambda: {Role.ALL: MapperKind.REPLICATED} + assert manager.get_replicated_roles() == frozenset() diff --git a/tests/unittest/_torch/executor/test_mooncake_store_common.py b/tests/unittest/_torch/executor/test_mooncake_store_common.py index d683c799ebbf..f1a4277e83d4 100644 --- a/tests/unittest/_torch/executor/test_mooncake_store_common.py +++ b/tests/unittest/_torch/executor/test_mooncake_store_common.py @@ -23,6 +23,7 @@ """ import json +import re from types import SimpleNamespace import pytest @@ -33,9 +34,11 @@ StoreRole, ) from tensorrt_llm._torch.pyexecutor.connectors.mooncake_store.keys import ( + REPLICATED_SHARD_KEY, BlockHashChain, KeyNamespace, ReuseScope, + sharded_shard_key, ) from tensorrt_llm._torch.pyexecutor.connectors.mooncake_store.staging import ( HostStagingPool, @@ -191,8 +194,7 @@ def test_key_namespace_separates_every_dimension(): base = dict( namespace="trtllm", model_key="m", - rank=0, - world_size=2, + shard_key=sharded_shard_key(0, 2), layer_group_id=0, tokens_per_block=32, bytes_per_page=1024, @@ -202,8 +204,7 @@ def test_key_namespace_separates_every_dimension(): for field, value in [ ("namespace", "other"), ("model_key", "n"), - ("rank", 1), - ("world_size", 4), + ("shard_key", sharded_shard_key(1, 2)), ("layer_group_id", 1), ("tokens_per_block", 64), ("bytes_per_page", 2048), @@ -211,6 +212,19 @@ def test_key_namespace_separates_every_dimension(): assert KeyNamespace(**{**base, field: value}).key(block_hash) != reference +def test_shard_key_distinguishes_rank_and_world_size(): + # Rank 3 of 8 holds different heads than rank 3 of 4, so both components + # have to appear. + assert sharded_shard_key(3, 8) != sharded_shard_key(3, 4) + assert sharded_shard_key(0, 2) != sharded_shard_key(1, 2) + + +def test_replicated_shard_key_cannot_collide_with_a_sharded_one(): + # A sharded component is always wr; the replicated one is a + # literal, so no rank/world-size pair can produce it. + assert not re.fullmatch(r"w\d+r\d+", REPLICATED_SHARD_KEY) + + # ---- config ---- diff --git a/tests/unittest/_torch/executor/test_mooncake_store_connector.py b/tests/unittest/_torch/executor/test_mooncake_store_connector.py index 2beea0e0f728..87927d260b59 100644 --- a/tests/unittest/_torch/executor/test_mooncake_store_connector.py +++ b/tests/unittest/_torch/executor/test_mooncake_store_connector.py @@ -62,7 +62,10 @@ format_diagnosis, reservation_start, ) -from tensorrt_llm._torch.pyexecutor.connectors.mooncake_store.keys import BlockHashChain +from tensorrt_llm._torch.pyexecutor.connectors.mooncake_store.keys import ( + REPLICATED_SHARD_KEY, + BlockHashChain, +) from tensorrt_llm._torch.pyexecutor.connectors.mooncake_store.metadata import ( PageTransfer, RequestTransfers, @@ -120,31 +123,37 @@ def close(self): self.closed = True -def make_layout(*, num_groups=1, regions_per_group=1, num_slots=8): - """A layout whose regions are laid out back to back in a fake address space.""" +def make_layout(*, num_groups=1, regions_per_group=1, num_slots=8, replicated_per_group=0): + """A layout whose regions are laid out back to back in a fake address space. + + `replicated_per_group` adds that many index-K style regions per group, + standing in for a role the manager declared replicated. + """ groups = [] base = 0x1000 + + def region(size, role, layer_id): + nonlocal base + built = KvCacheRegion( + base=base, + size=size, + stride=size, + num_slots=num_slots, + buffers=(KvCacheBufferRef(layer_id=layer_id, role=role),), + ) + base += size * num_slots + return built + for group_id in range(num_groups): - regions = [] - for region_id in range(regions_per_group): - size = 64 * (region_id + 1) - stride = size - regions.append( - KvCacheRegion( - base=base, - size=size, - stride=stride, - num_slots=num_slots, - buffers=(KvCacheBufferRef(layer_id=group_id, role="key"),), - ) - ) - base += stride * num_slots + regions = [region(64 * (i + 1), "key", group_id) for i in range(regions_per_group)] + replicated = [region(32, "index_key", group_id) for _ in range(replicated_per_group)] groups.append( KvCacheLayerGroupLayout( layer_group_id=group_id, layer_ids=(group_id,), window_size=None, regions=tuple(regions), + replicated_regions=tuple(replicated), ) ) return KvCacheLayout(tokens_per_block=TOKENS_PER_BLOCK, groups=tuple(groups)) @@ -553,6 +562,65 @@ def test_page_addressing_rejects_mixed_slot_counts(): PageAddressing(layout) +def test_page_addressing_rejects_a_replicated_region_with_its_own_slot_count(): + # Both classes are indexed by the group's one page-index space, so a + # replicated region that disagrees would silently address the wrong slot. + layout = KvCacheLayout( + tokens_per_block=TOKENS_PER_BLOCK, + groups=( + KvCacheLayerGroupLayout( + layer_group_id=0, + layer_ids=(0,), + window_size=None, + regions=(KvCacheRegion(base=0, size=8, stride=8, num_slots=4, buffers=()),), + replicated_regions=( + KvCacheRegion(base=64, size=8, stride=8, num_slots=8, buffers=()), + ), + ), + ), + ) + with pytest.raises(ValueError, match="slot counts"): + PageAddressing(layout) + + +def test_page_addressing_resolves_the_two_classes_separately(): + layout = make_layout(regions_per_group=2, replicated_per_group=2, num_slots=4) + addressing = PageAddressing(layout) + group = layout.groups[0] + + addresses, sizes = addressing.buffers(0, 2) + assert addresses == [region.base + region.stride * 2 for region in group.regions] + assert sizes == [region.size for region in group.regions] + + rep_addresses, rep_sizes = addressing.replicated_buffers(0, 2) + assert rep_addresses == [region.base + region.stride * 2 for region in group.replicated_regions] + assert rep_sizes == [region.size for region in group.replicated_regions] + + # The two payloads are disjoint, so their sizes do not double-count. + assert addressing.has_replicated(0) + assert addressing.bytes_per_page(0) == sum(region.size for region in group.regions) + assert addressing.replicated_bytes_per_page(0) == sum( + region.size for region in group.replicated_regions + ) + + +def test_page_addressing_reports_no_replicated_page_without_the_role(): + addressing = PageAddressing(make_layout(regions_per_group=2)) + assert not addressing.has_replicated(0) + assert addressing.replicated_bytes_per_page(0) == 0 + assert addressing.replicated_buffers(0, 0) == ([], []) + + +def test_page_addressing_registers_replicated_regions_too(): + # Unregistered memory cannot be RDMA'd, so omitting the replicated span + # would fail every transfer that touches it. + layout = make_layout(regions_per_group=1, replicated_per_group=1, num_slots=4) + replicated = layout.groups[0].replicated_regions[0] + span_end = replicated.base + replicated.stride * (replicated.num_slots - 1) + replicated.size + covered = PageAddressing(layout).registration_ranges() + assert any(start <= replicated.base and span_end <= end for start, end in covered) + + # ---- GPU pool mapping size ---- @@ -798,6 +866,107 @@ def test_worker_load_addresses_the_requested_page(store_config, fake_store): assert sizes == [expected_sizes] +def test_worker_names_the_replicated_page_without_a_rank( + store_config: Path, fake_store: FakeStore, monkeypatch: pytest.MonkeyPatch +) -> None: + """Index-K carries one key across the whole TP group.""" + layout = make_layout(replicated_per_group=1) + block_hash = b"\x03" * 16 + monkeypatch.setattr(worker_module, "mpi_world_size", lambda: 8) + + monkeypatch.setattr(worker_module, "mpi_rank", lambda: 0) + with make_worker(fake_store, layout=layout) as first: + first_sharded = first._namespaces[0].key(block_hash) + first_replicated = first._replicated_namespaces[0].key(block_hash) + monkeypatch.setattr(worker_module, "mpi_rank", lambda: 7) + with make_worker(fake_store, layout=layout) as last: + assert last._namespaces[0].key(block_hash) != first_sharded + assert last._replicated_namespaces[0].key(block_hash) == first_replicated + assert REPLICATED_SHARD_KEY in first_replicated + + +def test_worker_has_no_replicated_namespace_without_the_role(store_config, fake_store): + with make_worker(fake_store, layout=make_layout(num_groups=2)) as worker: + assert worker._replicated_namespaces == {} + # And nothing resolves, so the replicated path contributes no keys. + transfers = [RequestTransfers(1, [PageTransfer(b"\x00" * 16, 0, 0)])] + assert worker._resolve(transfers, replicated=True)[0] == [] + + +def test_worker_saves_and_loads_both_classes(store_config, fake_store): + layout = make_layout(replicated_per_group=1) + addressing = PageAddressing(layout) + block_hash = b"\x04" * 16 + with make_worker(fake_store, layout=layout) as worker: + worker._put([RequestTransfers(1, [PageTransfer(block_hash, 0, 2)])]) + + keys, addresses, sizes = fake_store.put_calls[0] + assert keys == [ + worker._namespaces[0].key(block_hash), + worker._replicated_namespaces[0].key(block_hash), + ] + assert addresses == [addressing.buffers(0, 2)[0], addressing.replicated_buffers(0, 2)[0]] + assert sizes == [addressing.buffers(0, 2)[1], addressing.replicated_buffers(0, 2)[1]] + + transfers = RequestTransfers(2, [PageTransfer(block_hash, 0, 5)]) + worker.bind_connector_meta(SimpleNamespace(loads=[transfers], saves=[])) + worker.start_load_kv(None) + loaded_keys, loaded_addresses, _ = fake_store.get_calls[-1] + assert loaded_keys == keys + assert loaded_addresses == [ + addressing.buffers(0, 5)[0], + addressing.replicated_buffers(0, 5)[0], + ] + + +def test_worker_saves_the_replicated_page_from_one_rank_only( + store_config: Path, fake_store: FakeStore, monkeypatch: pytest.MonkeyPatch +) -> None: + layout = make_layout(replicated_per_group=1) + block_hash = b"\x05" * 16 + monkeypatch.setattr(worker_module, "mpi_world_size", lambda: 4) + monkeypatch.setattr(worker_module, "mpi_rank", lambda: 2) + with make_worker(fake_store, layout=layout) as worker: + replicated_key = worker._replicated_namespaces[0].key(block_hash) + worker._put([RequestTransfers(1, [PageTransfer(block_hash, 0, 0)])]) + # Its own shard still goes out; the shared copy is rank 0's to write. + assert fake_store.put_calls[0][0] == [worker._namespaces[0].key(block_hash)] + + # But it still reads the shared copy, since it needs its own GPU copy. + fake_store.objects.add(replicated_key) + fake_store.objects.add(worker._namespaces[0].key(block_hash)) + transfers = RequestTransfers(2, [PageTransfer(block_hash, 0, 1)]) + worker.bind_connector_meta(SimpleNamespace(loads=[transfers], saves=[])) + worker.start_load_kv(None) + assert replicated_key in fake_store.get_calls[-1][0] + + +def test_adp_owners_all_save_the_replicated_page( + store_config: Path, fake_store: FakeStore, monkeypatch: pytest.MonkeyPatch +) -> None: + """ADP owners serve different requests, so gating them on rank would drop pages.""" + monkeypatch.setattr(worker_module, "mpi_world_size", lambda: 4) + monkeypatch.setattr(worker_module, "mpi_rank", lambda: 3) + layout = make_layout(replicated_per_group=1) + block_hash = b"\x06" * 16 + with make_worker(fake_store, layout=layout, enable_attention_dp=True) as worker: + worker._put([RequestTransfers(1, [PageTransfer(block_hash, 0, 0)])]) + assert worker._replicated_namespaces[0].key(block_hash) in fake_store.put_calls[0][0] + + +def test_worker_prefix_hit_needs_the_replicated_page(store_config, fake_store): + layout = make_layout(replicated_per_group=1) + with make_worker(fake_store, layout=layout) as worker: + block_hash = b"\x07" * 16 + fake_store.objects.add(worker._namespaces[0].key(block_hash)) + # The shard's own page landed but rank 0's shared write did not, so the + # block is unusable rather than half-loadable. + assert worker.count_prefix_hit([block_hash]) == 0 + + fake_store.objects.add(worker._replicated_namespaces[0].key(block_hash)) + assert worker.count_prefix_hit([block_hash]) == 1 + + def test_worker_reports_a_request_finished_once_its_saves_drain(store_config, fake_store): with make_worker(fake_store, layout=make_layout()) as worker: # One submission outstanding: the request is closed but must not be released. diff --git a/tests/unittest/disaggregated/test_minimax_m3_kv_transfer.py b/tests/unittest/disaggregated/test_minimax_m3_kv_transfer.py index 7ba51ec171e5..06529678968f 100644 --- a/tests/unittest/disaggregated/test_minimax_m3_kv_transfer.py +++ b/tests/unittest/disaggregated/test_minimax_m3_kv_transfer.py @@ -123,6 +123,9 @@ def fake_base_init(self, *args, **kwargs): Role.ALL: expected_main_mapper, Role.INDEX_KEY: MapperKind.REPLICATED, } + # Whichever layout the main K/V uses, index-K stays the one role a KV + # connector may store once for the whole TP group. + assert manager.get_replicated_roles() == frozenset({Role.INDEX_KEY}) def test_minimax_disagg_rejects_unmanaged_index_value(monkeypatch) -> None: