From 0972960b5237880fbe528f78632476858fbbba60 Mon Sep 17 00:00:00 2001 From: Balaram Buddharaju <169953907+brb-nv@users.noreply.github.com> Date: Fri, 25 Sep 2026 14:53:57 -0700 Subject: [PATCH 1/2] save initial changes Signed-off-by: Balaram Buddharaju <169953907+brb-nv@users.noreply.github.com> --- docs/source/features/kv-cache-connector.md | 4 +- .../pyexecutor/connectors/kv_cache_layout.py | 96 +++++++-- .../connectors/mooncake_store/addressing.py | 107 +++++++-- .../connectors/mooncake_store/keys.py | 35 ++- .../connectors/mooncake_store/worker.py | 122 +++++++++-- .../_torch/pyexecutor/kv_cache_manager_v2.py | 19 ++ .../_torch/executor/test_kv_cache_layout.py | 81 ++++++- .../executor/test_kv_cache_manager_v2.py | 36 ++++ .../executor/test_mooncake_store_common.py | 22 +- .../executor/test_mooncake_store_connector.py | 203 ++++++++++++++++-- .../test_minimax_m3_kv_transfer.py | 3 + 11 files changed, 644 insertions(+), 84 deletions(-) diff --git a/docs/source/features/kv-cache-connector.md b/docs/source/features/kv-cache-connector.md index 7425e4b2db1e..123a2bcdbd32 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..4fe990935749 100644 --- a/tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_layout.py +++ b/tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_layout.py @@ -31,6 +31,14 @@ 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 -- +MiniMax-M3's index-K is computed from a replicated projection. The manager +declares which roles those are, and they are described as ``replicated_regions`` +so a connector can store them once for the whole TP group rather than once per +rank. V2 may interleave the two classes in one pool, so the split is done by +aggregating each class separately rather than by slicing a merged range. """ from dataclasses import dataclass @@ -120,13 +128,23 @@ 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 declared + #: by the manager's ``get_replicated_roles()``. Empty for models that + #: register no such role, which is every model without a side cache. + 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 +205,43 @@ 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 lets shard-specific and replicated roles be described + separately even when V2 packed them into one interleaved pool. + """ + 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 +249,12 @@ 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. Their bytes are identical on every attention shard, so a connector can + address them once rather than once per rank. A manager that declares none + yields the same regions it would have without the split, since the buffer + set handed to the aggregator is then unchanged. """ impl = manager.impl init_config = impl.init_config @@ -210,6 +271,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) + # Role names rather than DataRole values: a region records the manager's + # native role string, and comparing strings keeps this free of the + # disaggregation types the manager uses to express the same 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 +286,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..80e9a2b050da 100644 --- a/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/addressing.py +++ b/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/addressing.py @@ -23,6 +23,12 @@ 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. Most roles hold bytes belonging +to a single attention shard; roles the manager declares replicated hold bytes +identical on every shard. They are addressed separately so the connector can +key the replicated page once for the whole TP group. Groups with no replicated +role report an empty second page, which the worker skips. """ from typing import Dict, Iterable, List, Optional, Sequence, Tuple @@ -90,9 +96,10 @@ def merge_intervals(intervals: Iterable[Tuple[int, int]]) -> List[Tuple[int, int """Collapse `(start, end)` byte ranges into a minimal disjoint cover. 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. + inside one pool allocation: sliding-window layer groups share it, and + separating replicated roles from shard-specific ones (MiniMax-M3's index-K + beside K/V) splits one pool into interleaved regions of both classes. + 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 +118,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 +132,26 @@ 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 are indexed by that same 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 +177,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 +196,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 +225,45 @@ 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. + + These bytes are identical on every attention shard, so the same page is + described here on every rank even though each rank names its own copy. + 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 +273,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 +286,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: replicated bytes are transferred like any other, so + leaving them unregistered would fail every transfer that touches them. 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 +301,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 +329,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..622bca15fc5e 100644 --- a/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/keys.py +++ b/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/keys.py @@ -23,6 +23,11 @@ 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 is where replicated roles pay off. A page whose bytes depend +on the shard is named per rank, but one whose bytes are identical on every rank +is named `REPLICATED_SHARD_KEY` instead, so a TP group stores it once rather +than once per rank. """ import hashlib @@ -34,6 +39,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 +146,33 @@ 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 rather than a rank, so a replicated page has one key for the whole +#: TP group. It 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: `sharded_shard_key(...)` for a page whose + #: content depends on the shard that produced it, `REPLICATED_SHARD_KEY` + #: for one whose content is identical on every shard. + shard_key: str layer_group_id: int tokens_per_block: int bytes_per_page: int @@ -159,7 +182,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..f978602e30a7 100644 --- a/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/worker.py +++ b/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/worker.py @@ -32,6 +32,11 @@ 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. Its shard-specific bytes are keyed per rank as +usual, while roles the manager declares replicated are keyed once for the whole +TP group, since every rank holds the same bytes. Only one rank writes that copy; +all of them read it, because each still needs the bytes in its own GPU memory. """ import os @@ -52,7 +57,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 +198,10 @@ 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] = {} + # Namespaces for roles whose bytes are identical on every shard, keyed + # without a rank so the whole TP group shares one copy. Present only + # for layer groups that hold such a 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 +352,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 +400,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 classes pass through the same slots, so a slot has to hold the + # larger of the two; a replicated page is usually the smaller one. 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 +438,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 +460,29 @@ 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. + + Replicated pages carry one key for the whole TP group, so every rank + holding identical bytes would otherwise race to write the same value. + Naming a single owner keeps that to one write. + + 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 +497,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 rather than one per shard, 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 +540,17 @@ def start_load_kv(self, stream: torch.cuda.Stream): self._reraise_save_error() keys, addresses, sizes, total_pages = self._resolve(metadata.loads) + # Every rank reads the replicated pages even though one rank wrote + # them: the bytes are shared only in the store, and each rank still + # needs its own GPU copy. Appending rather than loading separately + # keeps both 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 +713,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 +768,21 @@ 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. Layer groups holding no replicated + role contribute nothing, so the result is empty for a model + that declares none. + + 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 +791,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..d4e56196d8c7 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py @@ -1921,6 +1921,25 @@ 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 sharded roles rather than a role of its own, so it + is never replicated even when a subclass maps it to a layout kind. + + A KV connector uses this to give replicated bytes a single + rank-independent store key instead of one identical copy per rank. + Entries are inert unless the manager actually registers buffers for + that role, so the default is safe for managers that register none. + """ + 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..7e7c2a01cbac 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, not part of the + # shard-specific one: 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,68 @@ 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 splitting by class must leave this layout exactly + # as it was before the split existed. + 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 split must still separate the classes, which is + # what lets a connector store the replicated bytes once per TP group + # rather than once per rank. + 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..ac9ea772a6a6 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,39 @@ 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 without + # touching the connectors that consume this. + 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..fbf46f200082 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..09cd7e71f730 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 identical on every shard. + """ 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 overlap-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: + """The point of the split: one key for index-K 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 neither path gains a key. + 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: From ace7552449a1eb779ddd4b79deca4172a2feb84f Mon Sep 17 00:00:00 2001 From: Balaram Buddharaju <169953907+brb-nv@users.noreply.github.com> Date: Mon, 28 Sep 2026 22:32:16 -0700 Subject: [PATCH 2/2] [None][chore] Tighten comments and docstrings for replicated KV regions Signed-off-by: Balaram Buddharaju <169953907+brb-nv@users.noreply.github.com> --- docs/source/features/kv-cache-connector.md | 2 +- .../pyexecutor/connectors/kv_cache_layout.py | 37 ++++++++----------- .../connectors/mooncake_store/addressing.py | 29 ++++++--------- .../connectors/mooncake_store/keys.py | 15 +++----- .../connectors/mooncake_store/worker.py | 36 ++++++++---------- .../_torch/pyexecutor/kv_cache_manager_v2.py | 13 +++---- .../_torch/executor/test_kv_cache_layout.py | 11 ++---- .../executor/test_kv_cache_manager_v2.py | 3 +- .../executor/test_mooncake_store_common.py | 2 +- .../executor/test_mooncake_store_connector.py | 8 ++-- 10 files changed, 66 insertions(+), 90 deletions(-) diff --git a/docs/source/features/kv-cache-connector.md b/docs/source/features/kv-cache-connector.md index 123a2bcdbd32..bf7e24448365 100644 --- a/docs/source/features/kv-cache-connector.md +++ b/docs/source/features/kv-cache-connector.md @@ -73,7 +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. + * **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. diff --git a/tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_layout.py b/tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_layout.py index 4fe990935749..c27af2e24ff0 100644 --- a/tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_layout.py +++ b/tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_layout.py @@ -33,12 +33,11 @@ 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 -- -MiniMax-M3's index-K is computed from a replicated projection. The manager -declares which roles those are, and they are described as ``replicated_regions`` -so a connector can store them once for the whole TP group rather than once per -rank. V2 may interleave the two classes in one pool, so the split is done by -aggregating each class separately rather than by slicing a merged range. +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 @@ -130,9 +129,8 @@ class KvCacheLayerGroupLayout: window_size: Optional[int] #: Regions holding bytes specific to this attention shard. regions: Tuple[KvCacheRegion, ...] - #: Regions holding bytes identical on every attention shard, as declared - #: by the manager's ``get_replicated_roles()``. Empty for models that - #: register no such role, which is every model without a side cache. + #: 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 @@ -212,12 +210,11 @@ def _regions_for( num_slots: int, global_by_local: Dict[int, int], ) -> Tuple[KvCacheRegion, ...]: - """Aggregate ``buffer_ids`` into the regions they occupy in one slot. + """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 lets shard-specific and replicated roles be described - separately even when V2 packed them into one interleaved pool. + 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): @@ -251,10 +248,8 @@ def build_kv_cache_layout_v2(manager: "KVCacheManagerV2") -> KvCacheLayout: order, kv factor, or the number of pools. Roles the manager declares replicated are described as a separate region - set. Their bytes are identical on every attention shard, so a connector can - address them once rather than once per rank. A manager that declares none - yields the same regions it would have without the split, since the buffer - set handed to the aggregator is then unchanged. + set. A manager that declares none leaves replicated_regions empty and every + buffer in regions. """ impl = manager.impl init_config = impl.init_config @@ -271,9 +266,9 @@ 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) - # Role names rather than DataRole values: a region records the manager's - # native role string, and comparing strings keeps this free of the - # disaggregation types the manager uses to express the same declaration. + # 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] = [] diff --git a/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/addressing.py b/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/addressing.py index 80e9a2b050da..c4615aa91105 100644 --- a/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/addressing.py +++ b/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/addressing.py @@ -24,11 +24,9 @@ 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. Most roles hold bytes belonging -to a single attention shard; roles the manager declares replicated hold bytes -identical on every shard. They are addressed separately so the connector can -key the replicated page once for the whole TP group. Groups with no replicated -role report an empty second page, which the worker skips. +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 @@ -96,10 +94,9 @@ def merge_intervals(intervals: Iterable[Tuple[int, int]]) -> List[Tuple[int, int """Collapse `(start, end)` byte ranges into a minimal disjoint cover. A range may not be registered twice, but several regions routinely live - inside one pool allocation: sliding-window layer groups share it, and - separating replicated roles from shard-specific ones (MiniMax-M3's index-K - beside K/V) splits one pool into interleaved regions of both classes. - Merging first spares the caller that distinction. + inside one pool allocation: sliding-window layer groups share it, and a + 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]] = [] @@ -147,8 +144,7 @@ def __init__(self, layout: KvCacheLayout): 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. - # Replicated regions are indexed by that same space, so they are - # held to it too. + # 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) } @@ -237,10 +233,9 @@ def replicated_buffers( ) -> Tuple[List[int], List[int]]: """Addresses and sizes of one replicated page, in concatenation order. - These bytes are identical on every attention shard, so the same page is - described here on every rank even though each rank names its own copy. - A region crossing a GPU pool mapping boundary is split the same way as - in `buffers`. + 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. @@ -287,8 +282,8 @@ def registration_ranges(self) -> List[Tuple[int, int]]: 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. Both region classes - are covered: replicated bytes are transferred like any other, so - leaving them unregistered would fail every transfer that touches them. + 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 diff --git a/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/keys.py b/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/keys.py index 622bca15fc5e..34e7788e137c 100644 --- a/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/keys.py +++ b/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/keys.py @@ -24,10 +24,9 @@ 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 is where replicated roles pay off. A page whose bytes depend -on the shard is named per rank, but one whose bytes are identical on every rank -is named `REPLICATED_SHARD_KEY` instead, so a TP group stores it once rather -than once per rank. +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 @@ -147,8 +146,7 @@ def extend(self, tokens: Sequence[int]) -> Sequence[bytes]: #: Shard component for pages whose bytes do not depend on the attention shard. -#: A literal rather than a rank, so a replicated page has one key for the whole -#: TP group. It cannot collide with a sharded component, which is always +#: A literal cannot collide with a sharded component, which is always #: `wr`. REPLICATED_SHARD_KEY = "replicated" @@ -169,9 +167,8 @@ class KeyNamespace: namespace: str model_key: str - #: Which shard's bytes these are: `sharded_shard_key(...)` for a page whose - #: content depends on the shard that produced it, `REPLICATED_SHARD_KEY` - #: for one whose content is identical on every shard. + #: 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 diff --git a/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/worker.py b/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/worker.py index f978602e30a7..3d05df28ebab 100644 --- a/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/worker.py +++ b/tensorrt_llm/_torch/pyexecutor/connectors/mooncake_store/worker.py @@ -33,10 +33,9 @@ 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. Its shard-specific bytes are keyed per rank as -usual, while roles the manager declares replicated are keyed once for the whole -TP group, since every rank holds the same bytes. Only one rank writes that copy; -all of them read it, because each still needs the bytes in its own GPU memory. +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 @@ -198,9 +197,8 @@ 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] = {} - # Namespaces for roles whose bytes are identical on every shard, keyed - # without a rank so the whole TP group shares one copy. Present only - # for layer groups that hold such a role. + # 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. @@ -400,8 +398,8 @@ 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 classes pass through the same slots, so a slot has to hold the - # larger of the two; a replicated page is usually the smaller one. + # Both page classes pass through the same slots, so a slot has to hold + # the larger of the two. max_bytes_per_page = max( max( addressing.bytes_per_page(layer_group_id), @@ -464,9 +462,8 @@ def is_registered(self) -> bool: def _owns_replicated_saves(self) -> bool: """Whether this rank writes the replicated pages of its group. - Replicated pages carry one key for the whole TP group, so every rank - holding identical bytes would otherwise race to write the same value. - Naming a single owner keeps that to one write. + 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 @@ -497,8 +494,8 @@ 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 rather than one per shard, so a block - # whose replicated page was dropped reads as a miss on every rank. + # 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() ) @@ -540,10 +537,8 @@ def start_load_kv(self, stream: torch.cuda.Stream): self._reraise_save_error() keys, addresses, sizes, total_pages = self._resolve(metadata.loads) - # Every rank reads the replicated pages even though one rank wrote - # them: the bytes are shared only in the store, and each rank still - # needs its own GPU copy. Appending rather than loading separately - # keeps both classes inside the same transfer batches. + # 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 ) @@ -775,9 +770,8 @@ def _resolve( 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. Layer groups holding no replicated - role contribute nothing, so the result is empty for a model - that declares none. + 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 diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py index d4e56196d8c7..9e925e82c175 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_manager_v2.py @@ -1924,15 +1924,14 @@ def get_disagg_role_mapper_kinds(self) -> dict[DataRole, MapperKind]: 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 sharded roles rather than a role of its own, so it - is never replicated even when a subclass maps it to a layout kind. + 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 instead of one identical copy per rank. - Entries are inert unless the manager actually registers buffers for - that role, so the default is safe for managers that register none. + rank-independent store key. A role with no registered buffers + contributes nothing, so declaring one is harmless. """ return frozenset( role diff --git a/tests/unittest/_torch/executor/test_kv_cache_layout.py b/tests/unittest/_torch/executor/test_kv_cache_layout.py index 7e7c2a01cbac..a6c8ecf20dac 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_layout.py +++ b/tests/unittest/_torch/executor/test_kv_cache_layout.py @@ -116,8 +116,8 @@ def test_the_two_page_classes_are_sized_independently(self): 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, not part of the - # shard-specific one: a connector keys and transfers them separately. + # 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) @@ -290,8 +290,7 @@ def test_uniform_model_yields_one_full_slot_region(self): ["key", "value"] * 4, ) # The base manager declares INDEX_KEY replicated but registers no - # such buffer, so splitting by class must leave this layout exactly - # as it was before the split existed. + # such buffer, so every region stays shard-specific. self.assertEqual(group.replicated_regions, ()) self.assertEqual(group.replicated_bytes_per_page, 0) finally: @@ -301,9 +300,7 @@ def test_uniform_model_yields_one_full_slot_region(self): 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 split must still separate the classes, which is - # what lets a connector store the replicated bytes once per TP group - # rather than once per rank. + # 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 = ( 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 ac9ea772a6a6..7b760f174da6 100644 --- a/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py +++ b/tests/unittest/_torch/executor/test_kv_cache_manager_v2.py @@ -566,8 +566,7 @@ def test_replicated_roles_follow_the_disagg_declaration(): manager = object.__new__(KVCacheManagerV2) assert manager.get_replicated_roles() == frozenset({Role.INDEX_KEY}) - # A subclass declaring a new replicated side cache is picked up without - # touching the connectors that consume this. + # 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, diff --git a/tests/unittest/_torch/executor/test_mooncake_store_common.py b/tests/unittest/_torch/executor/test_mooncake_store_common.py index fbf46f200082..f1a4277e83d4 100644 --- a/tests/unittest/_torch/executor/test_mooncake_store_common.py +++ b/tests/unittest/_torch/executor/test_mooncake_store_common.py @@ -213,7 +213,7 @@ def test_key_namespace_separates_every_dimension(): def test_shard_key_distinguishes_rank_and_world_size(): - # rank 3 of 8 holds different heads than rank 3 of 4, so both components + # 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) diff --git a/tests/unittest/_torch/executor/test_mooncake_store_connector.py b/tests/unittest/_torch/executor/test_mooncake_store_connector.py index 09cd7e71f730..87927d260b59 100644 --- a/tests/unittest/_torch/executor/test_mooncake_store_connector.py +++ b/tests/unittest/_torch/executor/test_mooncake_store_connector.py @@ -127,7 +127,7 @@ def make_layout(*, num_groups=1, regions_per_group=1, num_slots=8, replicated_pe """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 identical on every shard. + standing in for a role the manager declared replicated. """ groups = [] base = 0x1000 @@ -596,7 +596,7 @@ def test_page_addressing_resolves_the_two_classes_separately(): 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 overlap-count. + # 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( @@ -869,7 +869,7 @@ def test_worker_load_addresses_the_requested_page(store_config, fake_store): def test_worker_names_the_replicated_page_without_a_rank( store_config: Path, fake_store: FakeStore, monkeypatch: pytest.MonkeyPatch ) -> None: - """The point of the split: one key for index-K across the whole TP group.""" + """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) @@ -888,7 +888,7 @@ def test_worker_names_the_replicated_page_without_a_rank( 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 neither path gains a key. + # 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] == []