From e76dd0db4ef726449916ecac4ca90917997f8406 Mon Sep 17 00:00:00 2001 From: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> Date: Thu, 1 Oct 2026 10:53:40 +0000 Subject: [PATCH 1/2] [None][feat] Add KV cache lending for transfer backends Add request-scoped staging and in-place lenders for KVCacheManagerV2, with opaque content identities, layout derivation, registration holds, and lease lifecycle handling. Include manager lifecycle and shrink hooks, accumulated fetch readiness, and layout, public API, staging, in-place, and host-tier tests. Signed-off-by: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> --- .../kv_cache/kv_cache_manager_v2.py | 47 +- .../pyexecutor/kv_cache/sharing/__init__.py | 94 + .../pyexecutor/kv_cache/sharing/_identity.py | 97 + .../pyexecutor/kv_cache/sharing/_layout.py | 592 ++++ .../pyexecutor/kv_cache/sharing/_lender.py | 1407 +++++++++ .../pyexecutor/kv_cache/sharing/_manager.py | 191 ++ .../pyexecutor/kv_cache/sharing/_slots.py | 233 ++ .../pyexecutor/kv_cache/sharing/_types.py | 281 ++ .../integration/test_lists/test-db/l0_a10.yml | 1 + .../executor/kv_cache/sharing/conftest.py | 954 ++++++ .../kv_cache/sharing/test_host_tier.py | 867 ++++++ .../kv_cache/sharing/test_in_place_lender.py | 853 ++++++ .../kv_cache/sharing/test_layout_identity.py | 1261 ++++++++ .../test_layout_matches_native_page_table.py | 313 ++ .../executor/kv_cache/sharing/test_names.py | 187 ++ .../kv_cache/sharing/test_public_surface.py | 847 ++++++ .../executor/kv_cache/sharing/test_slots.py | 358 +++ .../kv_cache/sharing/test_staging_lender.py | 2594 +++++++++++++++++ .../executor/test_disagg_receive_ordering.py | 1 + 19 files changed, 11177 insertions(+), 1 deletion(-) create mode 100644 tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/__init__.py create mode 100644 tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_identity.py create mode 100644 tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_layout.py create mode 100644 tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_lender.py create mode 100644 tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_manager.py create mode 100644 tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_slots.py create mode 100644 tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_types.py create mode 100644 tests/unittest/_torch/executor/kv_cache/sharing/conftest.py create mode 100644 tests/unittest/_torch/executor/kv_cache/sharing/test_host_tier.py create mode 100644 tests/unittest/_torch/executor/kv_cache/sharing/test_in_place_lender.py create mode 100644 tests/unittest/_torch/executor/kv_cache/sharing/test_layout_identity.py create mode 100644 tests/unittest/_torch/executor/kv_cache/sharing/test_layout_matches_native_page_table.py create mode 100644 tests/unittest/_torch/executor/kv_cache/sharing/test_names.py create mode 100644 tests/unittest/_torch/executor/kv_cache/sharing/test_public_surface.py create mode 100644 tests/unittest/_torch/executor/kv_cache/sharing/test_slots.py create mode 100644 tests/unittest/_torch/executor/kv_cache/sharing/test_staging_lender.py diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py index ed0cc2a89afe..14ec980f2049 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py @@ -12,6 +12,7 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. +import functools import hashlib import math import os @@ -1146,6 +1147,9 @@ class KVCacheManagerV2(BaseResourceManager): # Read by KvCacheCreator: a subclass that overrides the context # commit/history protocol opts out of generic reuse-match backoff. _supports_reuse_match_backoff = True + # The lender sharing.attach_* installs, one per manager for its life; None without one. + # On the class so it exists without running __init__. + _sharing = None def __init__( self, @@ -3438,6 +3442,7 @@ def revert_allocate_generation(self, req: LlmRequest) -> None: f"{req.py_request_id} from {kv_cache.capacity} to " f"{reverted_cap}" ) + self._after_shrink(request_id, kv_cache) def revert_allocate_context(self, req: LlmRequest) -> bool: """Undo this iteration's context resize. False means the cache was dropped, @@ -3468,6 +3473,7 @@ def revert_allocate_context(self, req: LlmRequest) -> bool: f"request {req.py_request_id} from {kv_cache.capacity} " f"to {pre_cap}" ) + self._after_shrink(req.py_request_id, kv_cache) if pre_cap > 0: kv_cache.suspend() return True @@ -4327,6 +4333,7 @@ def _prepare_draft_resources(self, scheduled_batch: ScheduledRequests): f"{req.py_request_id}: could not resize to {capacity} tokens" f"{self._draft_pool_diagnostic()}" ) + self._after_shrink(req.py_request_id, kv_cache) for req in scheduled_batch.generation_requests: kv_cache = self._mirror_draft_kv_cache(req) @@ -5286,6 +5293,13 @@ def free_resources(self, request: LlmRequest, pin_on_release: bool = False): self.impl.clear_stats_excluded(request.py_request_id) return kv_cache.discard_pending_stats() + if self._sharing is not None and self._sharing._on_free( + request.py_request_id, + kv_cache, + functools.partial(self.impl.clear_stats_excluded, request.py_request_id), + ): + self._free_lent(request.py_request_id, kv_cache) + return kv_cache.close() self.impl.clear_stats_excluded(request.py_request_id) if request.py_request_id in self._early_freed_index_requests: @@ -5293,6 +5307,24 @@ def free_resources(self, request: LlmRequest, pin_on_release: bool = False): else: self.index_mapper.remove_sequence(request.py_request_id) + def _after_shrink(self, request_id: int, kv_cache) -> None: + """Tell the lender, if one is attached, that ``kv_cache`` may have shrunk in place: blocks + past its capacity lost their pages, and growing again does not bring their contents back.""" + if self._sharing is not None: + self._sharing._on_shrink(request_id, kv_cache) + + def _free_lent(self, request_id: int, kv_cache) -> None: + """Free a request whose cache is lent, leaving the cache open for the lender to close + after its last lease. The index slot is freed now unless an early free did, with the cache + detached from it first so the later close writes nothing into the slot's next owner.""" + if request_id in self._early_freed_index_requests: + self._early_freed_index_requests.discard(request_id) + return + for i in range(self.max_beam_width): + for pool_idx in range(self.num_pools): + kv_cache.set_base_page_index_buf(i, pool_idx, None) + self.index_mapper.remove_sequence(request_id) + def get_layer_page_index_scale(self, layer_idx: int) -> int: """Page-index scale of this layer's KV buffer. Layers in one pool can have different scales (e.g. different head_dim), so per-layer callers @@ -5537,6 +5569,7 @@ def check_invalid_values_in_kv_cache(self, fill_with_zero: bool = False) -> bool return bool(has_invalid_values) def shutdown(self): + keeping_lent = self._sharing is not None and self._keep_lent_until_exit() for kv_cache in self.kv_cache_map.values(): kv_cache.close() self.kv_cache_map.clear() @@ -5547,7 +5580,8 @@ def shutdown(self): # its plan, which mutates manager state. if self.conversation_manager is not None: self.conversation_manager.clear() - self.impl.shutdown() + if not keeping_lent: + self.impl.shutdown() # Shut the streaming event manager down last so removals emitted during # cache / impl teardown (via the radix tree's own event-manager # reference) are still flushed before the publisher stops. Do not null @@ -5556,6 +5590,15 @@ def shutdown(self): if isinstance(self.event_manager, StreamingKVCacheEventManager): self.event_manager.shutdown() + def _keep_lent_until_exit(self) -> bool: + """Keep every cache still lent in place, and the pools holding it, until the process exits. + Kept caches leave ``kv_cache_map`` so ``shutdown`` does not close them. Returns whether any + was kept; ``shutdown`` must then skip ``impl.shutdown()``.""" + kept = self._sharing._on_shutdown(self.impl) + for request_id in [rid for rid, kv_cache in self.kv_cache_map.items() if kv_cache in kept]: + del self.kv_cache_map[request_id] + return bool(kept) + def get_max_resource_count(self) -> int: # TODO: implement this return 1 @@ -5760,6 +5803,8 @@ def update_resources( f"to capacity {new_capacity} and history length " f"{history_length} tokens at generation update" ) + if new_capacity is not None: + self._after_shrink(req.py_request_id, kv_cache) self._allocated_draft_lens.pop(req.py_request_id, None) def copy_batch_block_offsets( diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/__init__.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/__init__.py new file mode 100644 index 000000000000..fce18a2a4020 --- /dev/null +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/__init__.py @@ -0,0 +1,94 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Lending a KV cache manager v2's blocks to transfer backends. This module is the whole API; +every other module in the package is private, and the attach functions load the implementation.""" + +import typing as _typing + +from ._types import ( + GroupRun, + InPlaceLender, + Lease, + Part, + PartsHold, + Readiness, + RegionView, + StagingLender, + StagingOptions, +) + +if _typing.TYPE_CHECKING: + from ..kv_cache_manager_v2 import KVCacheManagerV2 + +__all__ = [ + "GroupRun", + "InPlaceLender", + "Lease", + "Part", + "PartsHold", + "Readiness", + "RegionView", + "StagingLender", + "StagingOptions", + "attach_in_place", + "attach_staging", +] + + +def attach_staging( + manager: "KVCacheManagerV2", *, scope: bytes, staging: StagingOptions +) -> StagingLender: + """Attach the manager's one lender, which relays whole blocks through host staging. + + Args: + manager: A ``KVCacheManagerV2`` that commits blocks to its prefix-reuse tree: block reuse + on, and joint reuse for a draft manager, since a publish lends only committed blocks. + scope: Equal exactly where KV bytes mean the same: model, weights, numerics, attention; + at most 65535 bytes. + staging: The staging size, in whole fetches. + + Returns: + The lender, installed on ``manager`` for the manager's life. + + Raises: + TypeError: ``manager`` is not a ``KVCacheManagerV2`` or has no mapping, ``scope`` is not + ``bytes``, or ``staging`` is not ``StagingOptions``. + ValueError: A lender is attached already; the manager is context-parallel, holds recurrent + state or commits no blocks (block reuse off); ``scope`` is longer than 65535 bytes; + ``staging.max_bytes`` is below one fetch. + """ + from ._lender import attach_staging as _attach + + return _attach(manager, scope=scope, staging=staging) + + +def attach_in_place(manager: "KVCacheManagerV2") -> InPlaceLender: + """Attach the manager's one lender, which lends requests' own device pages in place. + + Args: + manager: A ``KVCacheManagerV2``. Caches on loan at the manager's shutdown, with its device + pools, stay until the process exits. + + Returns: + The lender, installed on ``manager`` for the manager's life. + + Raises: + TypeError: ``manager`` is not a ``KVCacheManagerV2`` or has no mapping. + ValueError: A lender is attached already, or the manager is context-parallel or holds + recurrent state. + """ + from ._lender import attach_in_place as _attach + + return _attach(manager) diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_identity.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_identity.py new file mode 100644 index 000000000000..4007466f9bdb --- /dev/null +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_identity.py @@ -0,0 +1,97 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Names of the rows and staging parts instances exchange; equal names mean interchangeable bytes. +The format is a contract between instances: changing it bumps ``KEY_FORMAT_VERSION``.""" + +from __future__ import annotations + +import hashlib +from typing import Iterable, List, Sequence, Tuple + +import numpy as np + +from ._layout import canonical_order + +KEY_FORMAT_VERSION = 2 +KEY_BYTES = 32 +_NAMESPACE_BYTES = 16 +_GROUP_BYTES = 6 # the layer group's canonical index, shard count and shard index, u16 each +NAME_BYTES = _NAMESPACE_BYTES + KEY_BYTES + _GROUP_BYTES +_PART_LAYOUT_HEX = 16 + + +def _u16(value: int) -> bytes: + return int(value).to_bytes(2, "big") + + +def namespace(scope: bytes, layout_id: bytes) -> bytes: + """The 16 bytes every name starts with, from the caller's ``scope`` and the block layout's + ``layout_id`` (each length-prefixed). ``ValueError`` for an empty ``layout_id`` or a ``scope`` + longer than 65535 bytes.""" + if not layout_id: + raise ValueError("a namespace without layout_id would mix incompatible bytes") + if len(scope) > 0xFFFF: + raise ValueError(f"scope is {len(scope)} bytes; at most 65535 fit its length prefix") + digest = hashlib.sha256(b"trtllm-kv-object" + _u16(KEY_FORMAT_VERSION)) + digest.update(_u16(len(scope)) + scope) + digest.update(_u16(len(layout_id)) + layout_id) + return digest.digest()[:_NAMESPACE_BYTES] + + +class Identity: + """Names of one manager's rows and staging parts. ``layers[g]`` holds layer group ``g``'s global + layer ids and ``shards[g]`` the ``(count, index)`` share of its content this rank holds.""" + + namespace: bytes + layout_id: bytes + canonical: List[int] + + def __init__( + self, + scope: bytes, + layout_id: bytes, + layers: Sequence[Sequence[int]], + shards: Sequence[Tuple[int, int]], + ) -> None: + if len(shards) != len(layers): + raise ValueError(f"{len(layers)} layer groups but {len(shards)} shards") + self.layout_id = bytes(layout_id) + self.namespace = namespace(bytes(scope), self.layout_id) + self.canonical = canonical_order(layers) + self._namespace = np.frombuffer(self.namespace, dtype=np.uint8) + self._groups = [] + for index, (count, share) in zip(self.canonical, shards): + if count < 1 or not 0 <= share < count: + raise ValueError(f"shard {share} of {count}") + suffix = _u16(index) + _u16(count) + _u16(share) + self._groups.append(np.frombuffer(suffix, dtype=np.uint8)) + + def names(self, layer_group: int, keys: np.ndarray) -> np.ndarray: + """``uint8 (n, NAME_BYTES)``: the namespace, each row's 32-byte block key, then the layer + group's canonical index, shard count and shard index as big-endian ``u16``.""" + if not 0 <= layer_group < len(self._groups): + raise ValueError(f"unknown layer group {layer_group}") + keys = np.ascontiguousarray(keys, dtype=np.uint8).reshape(-1, KEY_BYTES) + out = np.empty((keys.shape[0], NAME_BYTES), dtype=np.uint8) + out[:, :_NAMESPACE_BYTES] = self._namespace + out[:, _NAMESPACE_BYTES : _NAMESPACE_BYTES + KEY_BYTES] = keys + out[:, _NAMESPACE_BYTES + KEY_BYTES :] = self._groups[layer_group] + return out + + def part_name(self, layer_groups: Iterable[int]) -> str: + """The name of the staging part serving the local ``layer_groups``: equal on instances laid + out alike.""" + groups = sorted({self.canonical[int(g)] for g in layer_groups}) + return f"{self.layout_id.hex()[:_PART_LAYOUT_HEX]}:lg{'+'.join(str(g) for g in groups)}" diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_layout.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_layout.py new file mode 100644 index 000000000000..c3139b5f01f8 --- /dev/null +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_layout.py @@ -0,0 +1,592 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""The layout of a manager's blocks as the lender needs it, and ``layout_id``, the digest of how a +block's bytes are laid out. Read from the manager's public layout surface and declarations.""" + +from __future__ import annotations + +import enum +import hashlib +import json +import struct +from dataclasses import dataclass +from typing import Dict, List, Literal, Mapping, Optional, Sequence, Tuple + +import numpy as np + +from . import _manager + +# A new field left out at its default keeps the version; changing the encoding or the meaning of +# an existing field bumps it. +LAYOUT_ID_VERSION = 2 + +# The fallback entry of the manager's per-role declarations (``Role.ALL``). +_ALL_ROLE = "all" + + +class BufferMapper(enum.IntEnum): + """How a buffer's bytes split across tensor-parallel ranks; the values of the mapper kinds + managers declare per role, read by value.""" + + INDEXED = 0 # head-major K/V, the default: a head range is one contiguous range per buffer + REPLICATED = 1 # bytes identical on every TP rank + NHD = 2 # token-major [token, head, dim] K/V: a head range is a slice inside every token + SECTIONED = 3 # [Sec0|Sec1|...], each section sharded independently (recurrent conv state) + + +@dataclass(frozen=True) +class BufferGeometry: + """Resharding geometry of one buffer; all ``None`` means no head axis. ``section_bytes`` sum to + the buffer's size; ``bytes_per_head * num_heads`` equals it.""" + + section_bytes: Optional[Tuple[int, ...]] = None + bytes_per_head: Optional[int] = None + num_heads: Optional[int] = None + + def validate(self) -> BufferGeometry: + """A copy with ``section_bytes`` as ints; ``ValueError`` unless every size given is + positive.""" + sections = self.section_bytes + if sections is not None: + sections = tuple(int(s) for s in sections) + if any(s <= 0 for s in sections): + raise ValueError(f"section sizes must be positive, got {sections}") + for name in ("bytes_per_head", "num_heads"): + value = getattr(self, name) + if value is not None and int(value) <= 0: + raise ValueError(f"{name} must be positive, got {value}") + return BufferGeometry(sections, self.bytes_per_head, self.num_heads) + + def check_size(self, size: int) -> None: + """``ValueError`` unless this geometry describes a buffer of ``size`` bytes.""" + if self.section_bytes is not None and sum(self.section_bytes) != size: + raise ValueError(f"sections {self.section_bytes} do not sum to {size} bytes") + if self.bytes_per_head is not None and size % self.bytes_per_head: + raise ValueError( + f"{size} bytes is not a whole number of {self.bytes_per_head}-byte heads" + ) + if ( + self.bytes_per_head is not None + and self.num_heads is not None + and self.bytes_per_head * self.num_heads != size + ): + raise ValueError( + f"{self.num_heads} heads of {self.bytes_per_head} bytes is not {size} bytes" + ) + + +@dataclass(frozen=True) +class ShardDesc: + """Which share (``index`` of ``count``) of a layer group's content this rank holds; ``(1, 0)`` + is the whole content, which every rank names alike. ``ValueError`` for an index outside the + count.""" + + count: int = 1 + index: int = 0 + + def __post_init__(self): + if self.count < 1 or not 0 <= self.index < self.count: + raise ValueError(f"shard {self.index} of {self.count}") + + +@dataclass(frozen=True) +class LayerGroupDesc: + """One layer group: its kind, window and sink blocks (which blocks exist), its device pool + group, its global layer ids and the share of its content this rank holds.""" + + kind: Literal["attention", "state"] + window: Optional[int] + sink_blocks: int + pool_group: int + layers: Tuple[int, ...] + shard: ShardDesc = ShardDesc() + + +@dataclass(frozen=True) +class BufferDesc: + """One buffer of one layer (global id) in a page: ``size`` bytes at ``offset`` in the slot of + ``pool`` = (pool group, pool index), how its bytes split across ranks, whether peers transfer + it, and its sub-pages per block. ``ValueError`` for a non-positive size or expansion.""" + + layer: int + role: str + pool: Tuple[int, int] + offset: int + size: int + mapper: BufferMapper = BufferMapper.INDEXED + geometry: BufferGeometry = BufferGeometry() + # False for local-only buffers peers skip; they still ride along wherever whole slots are + # copied. + transfer: bool = True + # The buffer's own tokens per block is tokens_per_block / expansion; a consumer reading inside + # the buffer that does not support the value must reject the layout. + expansion: int = 1 + + def __post_init__(self): + object.__setattr__(self, "role", str(self.role)) + mapper = self.mapper + if isinstance(mapper, str): + mapper = BufferMapper[mapper.upper()] + object.__setattr__(self, "mapper", BufferMapper(mapper)) + if not isinstance(self.geometry, BufferGeometry): + raise ValueError(f"geometry must be a buffer geometry, got {self.geometry!r}") + if self.size <= 0 or self.offset < 0: + raise ValueError(f"buffer at offset {self.offset} of {self.size} bytes") + if int(self.expansion) < 1: + raise ValueError(f"expansion must be a positive integer, got {self.expansion}") + self.geometry.check_size(self.size) + + +@dataclass(frozen=True) +class Layout: + """How pages are built, without addresses or slot counts: ``pool_groups[g]`` holds the slot + widths of device pool group ``g`` in pool order. ``ValueError`` for an expansion that does not + divide ``tokens_per_block``.""" + + tokens_per_block: int + pool_groups: Mapping[int, Tuple[int, ...]] + layer_groups: Tuple[LayerGroupDesc, ...] + buffers: Tuple[BufferDesc, ...] + + def __post_init__(self): + pool_groups = { + int(g): tuple(int(w) for w in widths) for g, widths in dict(self.pool_groups).items() + } + object.__setattr__(self, "pool_groups", pool_groups) + for b in self.buffers: + if self.tokens_per_block % int(b.expansion): + raise ValueError( + f"buffer ({b.layer}, {b.role}): expansion {b.expansion} does not divide " + f"{self.tokens_per_block} tokens per block" + ) + + +# Enums encode as the member name, floats as "f64:" plus the 16 hex digits of their big-endian +# IEEE-754 binary64 bits; a None value drops an object's key and is null in an array. +def _canonical(value): + if isinstance(value, enum.Enum): + return value.name + if value is None or isinstance(value, (bool, str)): + return value + if isinstance(value, np.bool_): + return bool(value) + if isinstance(value, (int, np.integer)): + return int(value) + if isinstance(value, (float, np.floating)): + return "f64:" + struct.pack(">d", float(value)).hex() + if isinstance(value, Mapping): + out = {} + for key, item in value.items(): + if not isinstance(key, str): + raise TypeError(f"canonical objects have string keys, got {key!r}") + if item is not None: + out[key] = _canonical(item) + return out + if isinstance(value, (list, tuple)): + return [_canonical(item) for item in value] + raise TypeError(f"{type(value).__name__} has no canonical encoding") + + +def canonical_bytes(document: Mapping) -> bytes: + """``document`` as ASCII JSON, keys sorted, no whitespace; enums by name, floats as ``f64:`` and + their big-endian bits in hex, ``None`` values dropped from objects.""" + return json.dumps( + _canonical(document), + sort_keys=True, + separators=(",", ":"), + ensure_ascii=True, + allow_nan=False, + ).encode("ascii") + + +def canonical_order(layers: Sequence[Sequence[int]]) -> List[int]: + """Each local layer group's canonical index: groups ordered by their smallest global layer id, + which pipeline parallelism does not change; ties keep the local order.""" + order = sorted(range(len(layers)), key=lambda g: (min(layers[g], default=1 << 62), g)) + canonical = [0] * len(layers) + for index, local in enumerate(order): + canonical[local] = index + return canonical + + +def _geometry_document(geometry: BufferGeometry) -> dict: + return { + "section_bytes": None if geometry.section_bytes is None else list(geometry.section_bytes), + "bytes_per_head": geometry.bytes_per_head, + "num_heads": geometry.num_heads, + } + + +def _buffer_document(layout: Layout, b: BufferDesc) -> dict: + # A buffer's position in the page: the slot widths of the pools before it in its group, plus + # its offset inside the slot. + widths = layout.pool_groups.get(b.pool[0]) + if widths is None: + raise ValueError(f"buffer ({b.layer}, {b.role}) lies in pool {b.pool}, not in the layout") + doc = { + "layer": int(b.layer), + "role": b.role, + "offset": sum(widths[: b.pool[1]]) + int(b.offset), + "size": int(b.size), + } + # Fields at their defaults are left out, so a new defaulted field keeps existing ids. + if b.mapper is not BufferMapper.INDEXED: + doc["mapper"] = b.mapper.name + geometry = {k: v for k, v in _geometry_document(b.geometry).items() if v is not None} + if geometry: + doc["geometry"] = geometry + if not b.transfer: + doc["transfer"] = False + if int(b.expansion) != 1: + doc["expansion"] = int(b.expansion) + return doc + + +def layout_document(layout: Layout) -> dict: + """The document ``layout_id`` hashes: tokens per block, groups in canonical order (layers, slot + widths, kind), buffers by (layer, role). Addresses, windows, sinks, shards, local numbering and + element types stay out.""" + canonical = canonical_order([desc.layers for desc in layout.layer_groups]) + groups = [None] * len(layout.layer_groups) + for local, desc in enumerate(layout.layer_groups): + widths = layout.pool_groups.get(desc.pool_group) + if widths is None: + raise ValueError( + f"layer group {local} maps to pool group {desc.pool_group}, which is not in the " + "layout" + ) + entry = { + "layers": sorted(int(layer) for layer in desc.layers), + "slot_bytes": [int(w) for w in widths], + } + if desc.kind != "attention": + entry["kind"] = desc.kind + groups[canonical[local]] = entry + buffers = sorted( + (_buffer_document(layout, b) for b in layout.buffers), + key=lambda d: (d["layer"], d["role"]), + ) + for before, after in zip(buffers, buffers[1:]): + if (before["layer"], before["role"]) == (after["layer"], after["role"]): + raise ValueError(f"buffer ({before['layer']}, {before['role']}) appears twice") + return { + "v": LAYOUT_ID_VERSION, + "tokens_per_block": int(layout.tokens_per_block), + "groups": groups, + "buffers": buffers, + } + + +def layout_id(layout: Layout) -> bytes: + """SHA-256 of ``canonical_bytes(layout_document(layout))``.""" + return hashlib.sha256(canonical_bytes(layout_document(layout))).digest() + + +@dataclass(frozen=True) +class DevicePool: + """One device pool of pool group ``group``: slot ``s`` lives at ``base + s * slot_bytes``.""" + + group: int + index: int + base: int + slot_bytes: int + num_slots: int + + +@dataclass(frozen=True, eq=False) +class ManagerLayout: + """What the lender reads of one manager, as plain values that hold no manager: the address-free + ``layout``, facts per layer group (tuples indexed by the local layer group) and per device pool + group (mappings keyed by the pool group).""" + + layout: Layout + tokens_per_block: int + pool_group_of: Tuple[int, ...] + windows: Tuple[Optional[int], ...] + sink_blocks: Tuple[int, ...] + recurrent: Tuple[bool, ...] + layers: Tuple[Tuple[int, ...], ...] + shards: Tuple[Tuple[int, int], ...] + page_bytes: Mapping[int, int] + device_pools: Mapping[int, Tuple[DevicePool, ...]] + + @property + def num_layer_groups(self) -> int: + """The manager's layer groups.""" + return len(self.pool_group_of) + + @property + def pool_groups(self) -> Tuple[int, ...]: + """The device pool groups, sorted; staging parts follow this order.""" + return tuple(sorted(self.device_pools)) + + +def global_layer_ids(manager, internal_ids: Sequence[int]) -> List[int]: + """The global id of each internal layer id, the same under any pipeline split: the declared id; + else for virtual layers ``model_layer * number_of_attention_types + attention_type``; else the + model layer.""" + # One declared id per internal layer of each layer group, in order (FP4 MLA's tail layers). + declared = getattr(manager, "get_disagg_global_layer_ids", None) + if declared is not None: + table: Dict[int, int] = {} + for lg, layer_ids in enumerate(manager.impl.layer_grouping): + ids = [int(gid) for gid in declared(lg)] + if len(ids) != len(layer_ids): + raise ValueError( + f"layer group {lg}: {len(ids)} declared global ids for {len(layer_ids)} layers" + ) + table.update(zip((int(lid) for lid in layer_ids), ids)) + return [table[int(lid)] for lid in internal_ids] + virtual = _manager.virtual_layers(manager) + if virtual is None: + pp_layers = getattr(manager, "pp_layers", None) + if pp_layers is None: + return [int(lid) for lid in internal_ids] + return [int(pp_layers[int(lid)]) for lid in internal_ids] + # The number of types counts every member of the enum, so stages holding different attention + # types agree. + inverse, num_types = virtual + return [inverse[int(lid)][0] * num_types + inverse[int(lid)][1] for lid in internal_ids] + + +def _layer_configs(manager) -> Mapping[int, object]: + """Internal layer id -> layer config, looked up by ``layer_id``, not by list position.""" + config = getattr(manager, "kv_cache_manager_py_config", None) + if config is None: + config = manager.impl.init_config + out: Dict[int, object] = {} + for position, layer in enumerate(config.layers): + out[int(getattr(layer, "layer_id", position))] = layer + return out + + +def _declared(manager, getter: str) -> Dict[str, object]: + method = getattr(manager, getter, None) + if method is None: + return {} + return {str(role): value for role, value in dict(method()).items()} + + +def _expansions(layer_config, tokens_per_block: int) -> Dict[str, int]: + """Role -> sub-pages per block, from the layer's ``tokens_per_block_override``.""" + out = {} + for buffer in getattr(layer_config, "buffers", ()) or (): + override = getattr(buffer, "tokens_per_block_override", None) + if override is not None: + if tokens_per_block % int(override): + raise ValueError( + f"buffer {buffer.role}: {override} tokens per block does not divide " + f"{tokens_per_block}" + ) + out[str(buffer.role)] = tokens_per_block // int(override) + return out + + +class _Heads: + """Head counts of attention buffers: per rank from the manager, in total from its inputs.""" + + def __init__(self, manager): + mapping = getattr(manager, "mapping", None) + dp = bool(getattr(mapping, "enable_attention_dp", False)) + self.tp_size = 1 if mapping is None or dp else max(int(getattr(mapping, "tp_size", 1)), 1) + self.tp_rank = ( + 0 if self.tp_size == 1 else int(getattr(mapping, "tp_rank", 0)) % self.tp_size + ) + self._per_rank = list(getattr(manager, "num_kv_heads_per_layer", ()) or ()) + self._total = getattr(manager, "num_kv_heads", None) + self._virtual = _manager.virtual_layers(manager) is not None + self._pp_layers = getattr(manager, "pp_layers", None) + + def per_rank(self, internal: int) -> int: + if not self._per_rank: + return 0 + index = internal if internal < len(self._per_rank) else 0 + return int(self._per_rank[index] or 0) + + def total(self, internal: int) -> Optional[int]: + if isinstance(self._total, int): + return int(self._total) + if self._total is None or self._virtual or self._pp_layers is None: + return None + if internal >= len(self._pp_layers): + return None # an extra internal layer (FP4 MLA's tail) is no model layer + value = self._total[int(self._pp_layers[internal])] + return None if value is None else int(value) + + def attention_shard(self, internal: int) -> ShardDesc: + """The share of a head-sharded attention buffer this rank holds: ``T`` heads over ``tp`` + ranks, rank ``r`` holds ``ceil(T / tp)`` from ``r * T // tp``.""" + if self.tp_size == 1: + return ShardDesc() + heads, total = self.per_rank(internal), self.total(internal) + if heads <= 0 or total is None or total <= 0: + return ShardDesc(self.tp_size, self.tp_rank) + # total // heads distinct shares; fewer heads than ranks are repeated, so a single head is + # one share that every rank holds. + count = max(total // heads, 1) + index = min((self.tp_rank * total // self.tp_size) // heads, count - 1) + return ShardDesc(count, index) + + def state_shard(self) -> ShardDesc: + return ShardDesc(self.tp_size, self.tp_rank) if self.tp_size > 1 else ShardDesc() + + +def _buffer_geometry( + declared: object, mapper: BufferMapper, size: int, heads: int, state: bool +) -> BufferGeometry: + """A buffer's head geometry: declared sections or head size, completed per buffer. + ``ValueError`` if the declared or completed geometry has a size that is not positive.""" + declared = BufferGeometry( + getattr(declared, "section_bytes", None), getattr(declared, "bytes_per_head", None) + ).validate() + if mapper in (BufferMapper.REPLICATED, BufferMapper.SECTIONED): + return declared + bytes_per_head = declared.bytes_per_head + if bytes_per_head is None and not state and heads > 0 and size % heads == 0: + bytes_per_head = size // heads + if bytes_per_head is None or size % bytes_per_head: + return declared + return BufferGeometry(declared.section_bytes, bytes_per_head, size // bytes_per_head).validate() + + +# A layer group is whole when every transferred buffer holds the same bytes on every rank +# (replicated, or one KV head shared by all ranks); else it is the share this rank holds. +def _group_shard(shards: Sequence[ShardDesc], fallback: ShardDesc) -> ShardDesc: + parts = {s for s in shards if s.count != 1} + if not parts: + return ShardDesc() + return parts.pop() if len(parts) == 1 else fallback + + +def derive_layout(manager) -> ManagerLayout: + """The layout of a ``KVCacheManagerV2`` and its device pools. ``ValueError`` if its layer + groups, layer configs and declarations do not agree.""" + impl = manager.impl + tokens_per_block = int(manager.tokens_per_block) + layer_cfg = _layer_configs(manager) + grouping = [[int(lid) for lid in group] for group in impl.layer_grouping] + pg_of_lg = [int(x) for x in impl.get_life_cycle_pool_group_indices()] + if len(pg_of_lg) != len(grouping): + raise ValueError(f"{len(grouping)} layer groups but {len(pg_of_lg)} pool-group indices") + internal = sorted({lid for group in grouping for lid in group}) + global_of = dict(zip(internal, global_layer_ids(manager, internal))) + + state_groups = set() + for lg, layer_ids in enumerate(grouping): + if not layer_ids or layer_ids[0] not in layer_cfg: + raise ValueError(f"layer group {lg} has no layer config (layers {layer_ids})") + if type(layer_cfg[layer_ids[0]]).__name__.startswith("Ssm"): + state_groups.add(lg) + + # TODO: the base manager declares head-major K/V whichever attention backend writes its pages, + # so a manager paired with a token-major backend gets the head-major layout_id unless it + # declares its layout itself. + mappers = { + role: BufferMapper(int(kind)) + for role, kind in _declared(manager, "get_disagg_role_mapper_kinds").items() + } + default_mapper = mappers.get(_ALL_ROLE, BufferMapper.INDEXED) + role_layouts = _declared(manager, "get_disagg_role_layouts") + ignored_getter = getattr(manager, "get_disagg_ignored_roles", None) + ignored = frozenset(str(r) for r in (ignored_getter() if ignored_getter else ())) + heads = _Heads(manager) + expansions = {lid: _expansions(cfg, tokens_per_block) for lid, cfg in layer_cfg.items()} + + device_pools: Dict[int, Tuple[DevicePool, ...]] = {} + buffers: List[BufferDesc] = [] + shards: Dict[int, List[ShardDesc]] = {lg: [] for lg in range(len(grouping))} + for pg in impl.pool_group_descs: + g = int(pg.pool_group_index) + pools = tuple( + DevicePool( + g, int(p.pool_index), int(p.base_address), int(p.slot_bytes), int(pg.num_slots) + ) + for p in pg.pools + ) + device_pools[g] = tuple(sorted(pools, key=lambda p: p.index)) + # One variant per layer group drawing from this pool group; pool ``i`` of a slot holds the + # ``i``-th coalesced buffer, whose members sit back to back. + for variant in pg.slot_desc.variants: + lg = int(variant.layer_group_id) + state = lg in state_groups + for pool_index, coalesced in enumerate(variant.coalesced_buffers): + size = int(coalesced.single_buffer_size) + for j, buffer_id in enumerate(coalesced.buffer_ids): + lid = int(buffer_id.layer_id) + role = str(buffer_id.role) + mapper = mappers.get(role, default_mapper) + transfer = role not in ignored + buffers.append( + BufferDesc( + layer=global_of.get(lid, lid), + role=role, + pool=(g, pool_index), + offset=j * size, + size=size, + mapper=mapper, + geometry=_buffer_geometry( + role_layouts.get(role), mapper, size, heads.per_rank(lid), state + ), + transfer=transfer, + expansion=expansions.get(lid, {}).get(role, 1), + ) + ) + if not transfer or mapper is BufferMapper.REPLICATED: + continue + shards[lg].append(heads.state_shard() if state else heads.attention_shard(lid)) + + layer_groups: List[LayerGroupDesc] = [] + for lg, layer_ids in enumerate(grouping): + first = layer_cfg[layer_ids[0]] + is_state = lg in state_groups + window = None + sink_tokens = 0 + if not is_state: + window = getattr(first, "window_size", None) + window = None if window is None else int(window) + sink_tokens = int(getattr(first, "num_sink_tokens", None) or 0) + # Layer groups with equal slot sizes may share a pool group, each drawing its own slots. + g = pg_of_lg[lg] + if g not in device_pools: + raise ValueError(f"layer group {lg} maps to pool group {g}, which has no pools") + layer_groups.append( + LayerGroupDesc( + kind="state" if is_state else "attention", + window=window, + sink_blocks=-(-sink_tokens // tokens_per_block), + pool_group=g, + layers=tuple(sorted(global_of[lid] for lid in layer_ids)), + shard=_group_shard(shards[lg], ShardDesc(heads.tp_size, heads.tp_rank)), + ) + ) + + layout = Layout( + tokens_per_block=tokens_per_block, + pool_groups={g: tuple(p.slot_bytes for p in pools) for g, pools in device_pools.items()}, + layer_groups=tuple(layer_groups), + buffers=tuple(buffers), + ) + return ManagerLayout( + layout=layout, + tokens_per_block=tokens_per_block, + pool_group_of=tuple(pg_of_lg), + windows=tuple(desc.window for desc in layer_groups), + sink_blocks=tuple(desc.sink_blocks for desc in layer_groups), + recurrent=tuple(desc.kind == "state" for desc in layer_groups), + layers=tuple(desc.layers for desc in layer_groups), + shards=tuple((desc.shard.count, desc.shard.index) for desc in layer_groups), + page_bytes={g: sum(p.slot_bytes for p in pools) for g, pools in device_pools.items()}, + device_pools=device_pools, + ) diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_lender.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_lender.py new file mode 100644 index 000000000000..d00c1ef805d4 --- /dev/null +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_lender.py @@ -0,0 +1,1407 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""The staging and in-place lenders, their leases, and the hooks the manager calls on free and +shutdown. Like the manager, one thread at a time uses them: the builder, the executor loop, then +the shutdown thread. There are no locks and no background threads.""" + +from __future__ import annotations + +import collections +import traceback +import weakref +from dataclasses import dataclass, field +from typing import ( + TYPE_CHECKING, + Callable, + Deque, + Dict, + FrozenSet, + List, + Mapping, + Optional, + Sequence, + Set, + Tuple, +) + +import numpy as np +import torch +from cuda.bindings import driver as drv +from cuda.bindings import runtime as cudart + +from tensorrt_llm._utils import prefer_pinned +from tensorrt_llm.logger import logger + +from . import _manager +from ._identity import Identity +from ._layout import ManagerLayout, derive_layout, layout_id +from ._slots import Runs, Slots, slot_counts +from ._types import GroupRun, Part, Readiness, RegionView, StagingOptions + +if TYPE_CHECKING: + from ...llm_request import LlmRequest + +# Staging memory, caches and page-index buffers kept until the process exits; only a clean shutdown +# or a last loan's end removes an entry. One thread at a time changes it, so it takes no lock. +_KEPT: List[object] = [] + +_SHUT_DOWN = "the KV cache manager shut down" +_SUSPENDED = "the request's cache is suspended" +_SCRATCH = "the request's cache has SWA scratch reuse on, so its window blocks may not keep a fetch" + + +def _retained() -> Tuple[object, ...]: + """What is kept until exit, oldest first; for tests.""" + return tuple(_KEPT) + + +def _let_go(owner: object) -> None: + """Drop ``owner`` from the keep list, compared by identity.""" + for index in range(len(_KEPT) - 1, -1, -1): + if _KEPT[index] is owner: + del _KEPT[index] + return + + +class _HostMemory: + """Host memory of exactly ``nbytes``, page-locked where pinning pays off. Only ``free`` releases + it, once; nothing frees it on collection, so memory kept until exit stays mapped.""" + + def __init__(self, nbytes: int) -> None: + self.nbytes = nbytes + self._pageable: Optional[np.ndarray] = None + if prefer_pinned(): + # The driver pins the size asked; torch's pinned allocator rounds it to a power of two. + error, address = cudart.cudaHostAlloc(nbytes, cudart.cudaHostAllocDefault) + if error != cudart.cudaError_t.cudaSuccess: + raise MemoryError(f"pinning {nbytes} bytes of staging memory failed: {error}") + self.address = int(address) + else: + self._pageable = np.empty(nbytes, dtype=np.uint8) + self.address = int(self._pageable.ctypes.data) + + def free(self) -> None: + """Release the memory; later calls do nothing.""" + address, self.address = self.address, 0 + if not address or self._pageable is not None: + self._pageable = None + return + (error,) = cudart.cudaFreeHost(address) + if error != cudart.cudaError_t.cudaSuccess: + logger.warning(f"KV cache lender: freeing the staging memory failed: {error}") + + +def _allocate(nbytes: int) -> _HostMemory: + """One host allocation of ``nbytes`` for the staging parts, appended to the keep list.""" + # TODO: relay through the KV cache manager's own host tier once its host pools keep fixed, + # registrable addresses (resizing with mremap moves them); keep this staging for managers + # without a host tier. + memory = _HostMemory(max(int(nbytes), 1)) + _KEPT.append(memory) + return memory + + +def _context_parallel_size(manager) -> int: + """The manager's context-parallel size; ``TypeError`` for a manager without a mapping.""" + mapping = getattr(manager, "mapping", None) + if mapping is None: + raise TypeError("the KV cache lender needs a manager with a mapping") + return int(mapping.cp_size) + + +def _checked_layout(manager) -> ManagerLayout: + """The checks both attaches share: a v2 manager, no context parallelism, no lender yet, no + recurrent state.""" + _manager.require_v2(manager) + cp_size = _context_parallel_size(manager) + if cp_size > 1: + # The layout derives only tensor-parallel shards: context-parallel ranks would name their + # blocks alike while holding different pages. + raise ValueError( + f"context parallelism (cp_size={cp_size}) is not supported: its ranks would give " + "different pages the same names" + ) + # TODO: one lender per manager for now, so staging and in-place never attach together; a later + # internal owner could serve both views. + if _manager.attached(manager) is not None: + raise ValueError("a lender is already attached to this KV cache manager") + layout = derive_layout(manager) + recurrent = [lg for lg, state in enumerate(layout.recurrent) if state] + if recurrent: + raise ValueError( + f"layer groups {recurrent} hold recurrent state, which the lender does not lend" + ) + return layout + + +def _check_commits(manager) -> None: + """``ValueError`` for a manager that commits no blocks: a publish lends only committed ones.""" + if not _manager.commits_blocks(manager): + raise ValueError( + "staging needs a manager that commits blocks: block reuse on, and joint reuse for a " + "draft manager; with block reuse off nothing could ever be published" + ) + + +def attach_staging(manager, *, scope: bytes, staging: StagingOptions) -> Staging: + """The public ``attach_staging``: a ``Staging`` lender installed on ``manager``.""" + return _attach_staging(manager, scope=scope, staging=staging) + + +def _attach_staging( + manager, *, scope: bytes, staging: StagingOptions, cls: Optional[type] = None +) -> Staging: + """``attach_staging`` with the lender class as a parameter (``Staging`` when ``None``), so a + test can attach a subclass that breaks one rule.""" + layout = _checked_layout(manager) + _check_commits(manager) + if not isinstance(scope, bytes): + raise TypeError(f"scope must be bytes, got {type(scope).__name__}") + if not isinstance(staging, StagingOptions): + raise TypeError(f"staging must be StagingOptions, got {type(staging).__name__}") + counts = slot_counts(layout, staging) + identity = Identity(scope, layout_id(layout.layout), layout.layers, layout.shards) + slots = {g: int(counts.get(g, 0)) for g in layout.pool_groups} + served: Dict[int, List[int]] = {g: [] for g in layout.pool_groups} + for lg, g in enumerate(layout.pool_group_of): + served[int(g)].append(lg) + sizes = {g: slots[g] * int(layout.page_bytes[g]) for g in layout.pool_groups} + memory = _allocate(sum(sizes.values())) + base = memory.address + parts = [] + offset = 0 + for g in layout.pool_groups: + name = identity.part_name(served[g]) + parts.append(Part(name, base + offset, sizes[g], int(layout.page_bytes[g]), slots[g])) + offset += sizes[g] + lender = (cls or Staging)( + weakref.ref(manager), layout, identity, tuple(parts), Slots(slots), weakref.ref(memory) + ) + lender._keep_index_buffer(manager) + # Installed last, so a failure above leaves nothing attached. + _manager.install(manager, lender) + logger.info( + f"KV cache lender: namespace {identity.namespace.hex()}, staging {offset >> 20} MiB " + f"in {len(parts)} parts" + ) + return lender + + +def attach_in_place(manager) -> InPlace: + """The public ``attach_in_place``: an ``InPlace`` lender installed on ``manager``.""" + layout = _checked_layout(manager) + lender = InPlace(weakref.ref(manager), layout) + _manager.install(manager, lender) + return lender + + +class _Copy: + """Copies queued together, complete once their event says so. ``done`` asks the event at most + once per round of the lender's progress, and never again once it reported completion.""" + + def __init__(self, event: torch.cuda.Event) -> None: + self._event = event + self._done = False + self._round: Optional[int] = None + + def done(self, round_: Optional[int] = None) -> bool: + """Whether the copies have completed, without waiting; ``round_=None`` always asks.""" + if not self._done and (round_ is None or round_ != self._round): + self._round = round_ + self._done = bool(self._event.query()) + return self._done + + def wait(self) -> None: + """Block the host until the copies have completed.""" + if not self._done: + self._event.synchronize() + self._done = True + + +def _key_column(keys: List[bytes]) -> np.ndarray: + return np.frombuffer(b"".join(keys), dtype=np.uint8).reshape(len(keys), 32) + + +def _no_cache(request_id: int) -> str: + return f"request {request_id} has no KV cache" + + +def _stale(manager, layout: ManagerLayout, lg: int, history: int) -> Tuple[int, int]: + """Block ordinals ``[beg, end)`` behind layer group ``lg``'s window at ``history``.""" + if layout.windows[lg] is None: + return 0, 0 + return _manager.stale_blocks(manager, lg, history) + + +def _needed_ordinals( + manager, layout: ManagerLayout, lg: int, start_block: int, end_block: int, history: int +) -> np.ndarray: + """Ordinals of ``lg`` in ``[start_block, end_block)`` that a history of ``history`` tokens + still reads: all of them for full attention; the sinks and the window otherwise.""" + ordinals = np.arange(start_block, end_block, dtype=np.int64) + stale_beg, stale_end = _stale(manager, layout, lg, history) + return ordinals[(ordinals < stale_beg) | (ordinals >= stale_end)] + + +def _checked_masks(view: RegionView, masks: Sequence[np.ndarray]) -> List[np.ndarray]: + """Copies of ``masks``; ``ValueError`` unless one bool mask of shape ``(len(run),)`` per run.""" + masks = list(masks) + if len(masks) != len(view.runs): + raise ValueError(f"{len(masks)} masks for {len(view.runs)} runs") + out = [] + for run, mask in zip(view.runs, masks): + mask = np.asarray(mask) + if mask.dtype != np.bool_ or mask.shape != (len(run),): + raise ValueError( + f"layer group {run.layer_group}: the mask must be bool of shape ({len(run)},), " + f"got {mask.dtype} {mask.shape}" + ) + out.append(mask.copy()) + return out + + +@dataclass(eq=False) +class _Rows: + """Per layer group, in order: ordinals, device slots and their staging slots, aligned.""" + + layer_groups: List[int] + ordinals: List[np.ndarray] + device_slots: List[np.ndarray] + staging_slots: List[np.ndarray] = field(default_factory=list) + + @property + def num_rows(self) -> int: + return sum(len(o) for o in self.ordinals) + + +@dataclass(eq=False) +class _Fetch: + """One write lease's fetch into one cache: delivered once its marked rows' copy is queued + without error, settled once that copy has completed.""" + + start: int + end: int + kv: object + delivered: bool = False + copy: Optional[_Copy] = None + + +@dataclass(eq=False) +class _Delivered: + """What the fetches into one cache delivered: per layer group, by block ordinal, the rows whose + copy was queued and that no shrink freed since; ``origin`` is the lowest fetch start.""" + + kv: object + origin: int + blocks: List[np.ndarray] + usable: Optional[Tuple[int, int]] = None # (committed tokens, usable_until) last computed + + +class Staging: + """``StagingLender`` over one manager, which it references weakly. Every call first returns the + slots of settled leases and grants waiting ones in order, so progress needs no new traffic.""" + + def __init__( + self, + manager: weakref.ref, + layout: ManagerLayout, + identity: Identity, + parts: Tuple[Part, ...], + slots: Slots, + memory: weakref.ref, + ) -> None: + self._manager_ref = manager + self._layout = layout + self._identity = identity + self._parts = tuple(parts) + self._slots = slots + # Only the keep list holds the staging memory strongly; this finds it there at shutdown. + self._memory = memory + self._part_of_group = {g: i for i, g in enumerate(layout.pool_groups)} + self._any_window = any(w is not None for w in layout.windows) + self._fetches: Dict[int, _Fetch] = {} # the latest fetch into each request + self._delivered: Dict[int, _Delivered] = {} + # Requests whose history a fetch moved past the committed tokens: the cache it grew. + self._advanced: Dict[int, Tuple[object, int]] = {} + self._line: Deque[_StagingLease] = collections.deque() # waiting for slots, in order + self._holding: List[_StagingLease] = [] # granted, slots not yet returned + # Open leases, held strongly: an open lease keeps the staging memory at shutdown even + # once its holder has dropped it. + self._unreleased: Set[_StagingLease] = set() + # Open parts holds, held strongly too: a dropped hold still keeps the memory. + self._holds: Set[_PartsHold] = set() + self._quarantined: List[Runs] = [] # slots a failed copy may still touch; never reused + self._closed = False + self._index_buffer: Optional[object] = None # kept from the attach until the shutdown + self._round = 0 # rounds of progress, so each asks a copy's event at most once + + @property + def parts(self) -> Tuple[Part, ...]: + """See ``StagingLender.parts``.""" + return self._parts + + def hold_parts(self) -> _PartsHold: + """See ``StagingLender.hold_parts``.""" + if self._live() is None: + # The memory was freed or kept at the shutdown; a later hold changes neither. + return _PartsHold(None) + hold = _PartsHold(self) + self._holds.add(hold) + return hold + + def lend_read(self, request: LlmRequest, start: int, end: int) -> _StagingLease: + """See ``StagingLender.lend_read``.""" + start, end = self._whole_blocks(start, end) + request_id = int(request.py_request_id) + manager = self._live() + if manager is None: + return _StagingLease._failed(self, "read", request_id, _SHUT_DOWN) + self._progress() + kv = _manager.kv_of(manager, request_id) + if kv is None: + return _StagingLease._failed(self, "read", request_id, _no_cache(request_id)) + state = _manager.cache_state(kv) + if end > state.committed: + raise ValueError( + f"the range ends at {end}, past the {state.committed} committed tokens" + ) + if not state.active: + return _StagingLease._failed(self, "read", request_id, _SUSPENDED) + rows = self._lendable(manager, state.history, self._rows(manager, kv, start, end)) + counts = self._counts(rows.ordinals) + self._check_fits(counts) + keys = self._keys_for(manager, request, kv, rows.ordinals) + lease = _StagingLease(self, "read", request_id, kv, rows, keys) + self._open(lease, counts) + return lease + + def lend_write(self, request: LlmRequest, start: int, end: int) -> _StagingLease: + """See ``StagingLender.lend_write``.""" + start, end = self._whole_blocks(start, end) + request_id = int(request.py_request_id) + manager = self._live() + if manager is None: + return _StagingLease._failed(self, "write", request_id, _SHUT_DOWN) + self._progress() + layout = self._layout + tpb = int(layout.tokens_per_block) + # Rows and their names come from the layout for a history of ``end``, so every + # ValueError is raised before the cache changes. + ordinals = [ + _needed_ordinals(manager, layout, lg, start // tpb, end // tpb, end) + for lg in range(layout.num_layer_groups) + ] + counts = self._counts(ordinals) + self._check_fits(counts) + kv = _manager.kv_of(manager, request_id) + if kv is None: + return _StagingLease._failed(self, "write", request_id, _no_cache(request_id)) + state = _manager.cache_state(kv) + if start < (state.committed // tpb) * tpb: + raise ValueError( + f"the range starts at {start}, inside the committed whole blocks of " + f"{state.committed} tokens" + ) + previous = self._fetches.get(request_id) + if previous is not None and previous.kv is kv and not self._report_settled(previous): + raise ValueError(f"request {request_id} already has an unsettled fetch") + keys = self._keys_for(manager, request, kv, ordinals) + if self._any_window and end < state.history: + raise ValueError( + f"the range ends at {end}, below the history of {state.history} tokens its " + "windows keep" + ) + if not state.active: + return _StagingLease._failed(self, "write", request_id, _SUSPENDED) + if self._scratch_reuse_on(kv): + return _StagingLease._failed(self, "write", request_id, _SCRATCH) + # With a window the history moves to ``end``, so windows need pages only for the blocks a + # history of that length reads. + position = end if self._any_window else state.history + if not _manager.grow(manager, request, kv, position, end): + return _StagingLease._failed( + self, "write", request_id, f"no free pages to grow the cache to {end} tokens" + ) + # The cache has grown, and readiness accounts for it whatever this lease's outcome. + if position > state.committed: + self._advanced[request_id] = (kv, state.committed) + fetch = _Fetch(start, end, kv) + self._fetches[request_id] = fetch + slots = [] + doomed = None + for lg, lg_ordinals in enumerate(ordinals): + pages = _manager.pages(kv, lg) + lg_slots = np.full(len(lg_ordinals), -1, dtype=np.int64) + inside = lg_ordinals < len(pages) + lg_slots[inside] = pages[lg_ordinals[inside]] + if doomed is None and np.any(lg_slots < 0): + missing = lg_ordinals[lg_slots < 0].tolist() + doomed = f"layer group {lg}: blocks {missing[:8]} have no page" + slots.append(lg_slots) + rows = _Rows(list(range(layout.num_layer_groups)), ordinals, slots) + lease = _StagingLease(self, "write", request_id, kv, rows, keys, fetch) + if doomed is not None: + # It fails at its first poll, which abandons the fetch. + lease._doomed = doomed + self._unreleased.add(lease) + return lease + self._open(lease, counts) + return lease + + def readiness(self, request: LlmRequest) -> Optional[Readiness]: + """See ``StagingLender.readiness``.""" + request_id = int(request.py_request_id) + manager = self._live() + if manager is None: + raise ValueError(f"{_no_cache(request_id)}: {_SHUT_DOWN}") + self._progress() + kv = _manager.kv_of(manager, request_id) + if kv is None: + raise ValueError(_no_cache(request_id)) + state = _manager.cache_state(kv) + committed, history = state.committed, state.history + # A record of a replaced cache (a restart) is moot. + fetch = self._fetches.get(request_id) + if fetch is not None and fetch.kv is not kv: + del self._fetches[request_id] + fetch = None + delivered = self._delivered.get(request_id) + if delivered is not None and delivered.kv is not kv: + del self._delivered[request_id] + delivered = None + advanced = self._advanced.get(request_id) + if advanced is not None and advanced[0] is not kv: + del self._advanced[request_id] + advanced = None + if fetch is not None and not self._report_settled(fetch): + return None + if delivered is None: + if advanced is not None: + # Grown for a fetch that delivered nothing: nothing past committed is computed. + return Readiness(committed, history) + return Readiness(max(committed, history), history) + # A shrink the manager did not report still shows as blocks past the capacity. + self._void_past(delivered, kv) + if delivered.usable is None or delivered.usable[0] != committed: + delivered.usable = (committed, self._usable_until(manager, delivered, committed)) + usable = delivered.usable[1] + return Readiness(int(usable), int(self._floor(manager, history, delivered.origin))) + + def _on_free(self, request_id: int, kv_cache, after_close: Callable[[], None]) -> bool: + """Manager hook, after the request's cache left the map: its waiting leases fail and its + fetch records go. Returns ``False``: the manager closes the cache itself. Never raises.""" + try: + if self._live() is None: + return False + # Granted leases go on: a read's copy is queued already, and a write's marks copy + # nothing into pages other than the ones lent. + self._fail_waiting(int(request_id), "the request was freed") + self._fetches.pop(int(request_id), None) + self._delivered.pop(int(request_id), None) + self._advanced.pop(int(request_id), None) + self._progress() + except Exception: + # A hook never raises into the manager's free: the cache is then closed as usual, + # and granted leases keep their slots. + logger.error(f"KV cache lender: freeing request {request_id}: {traceback.format_exc()}") + return False + + def _on_shrink(self, request_id: int, kv_cache) -> None: + """Manager hook, right after the request's cache may have shrunk in place: delivered rows + past its blocks lost their pages for good. Never raises.""" + try: + delivered = self._delivered.get(int(request_id)) + if self._live() is not None and delivered is not None and delivered.kv is kv_cache: + self._void_past(delivered, kv_cache) + except Exception: + # A hook never raises into the manager's resize; readiness voids the rows it sees past + # the capacity. + logger.error( + f"KV cache lender: shrink of request {request_id}: {traceback.format_exc()}" + ) + + def _on_shutdown(self, impl) -> FrozenSet[object]: + """Manager hook, first in its shutdown, acting once: waits for its copies, fails waiting + leases, and frees the staging memory unless a lease or a hold is open or a slot was lost to + a failed copy. Returns ``frozenset()``; never raises.""" + if self._closed: + return frozenset() + try: + # The lender's only host wait: no more work comes on the stream. + for lease in self._holding: + if lease._copy is not None: + lease._copy.wait() + self._round += 1 + self._recycle() + for lease in list(self._line): + self._fail(lease, _SHUT_DOWN) + for lease in [lease for lease in self._unreleased if lease._doomed is not None]: + self._fail(lease, lease._doomed) + self._closed = True + self._let_go_index_buffer() + if not self._memory_in_use(): + # Every copy on it has completed above, and no lease or hold reaches it any more. + memory = self._memory() + _let_go(memory) + memory.free() + else: + logger.warning( + f"KV cache lender: keeping {sum(p.nbytes for p in self._parts) >> 20} MiB of " + f"staging memory until exit: {len(self._unreleased)} leases open, " + f"{len(self._holds)} parts holds open, " + f"{len(self._quarantined)} slot runs lost to failed copies" + ) + except Exception: + # A hook never raises into the manager's shutdown; the staging memory then stays in + # the keep list until exit. + self._closed = True + logger.error(f"KV cache lender: shutting down: {traceback.format_exc()}") + return frozenset() + + # One rule per method, so a test subclass that breaks exactly one rule overrides one method. + + def _progress(self) -> None: + """Return the slots of settled leases, then grant waiting leases in order.""" + # Progress happens only here, inside the lender's calls, which the executor loop makes + # every iteration through polls and readiness; no background thread drives it. + if self._live() is None: + return + self._round += 1 + self._recycle() + self._grant_waiting() + + def _landed(self, copy: Optional[_Copy]) -> bool: + """``copy`` has completed (or there is none), its event asked at most once this round.""" + return copy is None or copy.done(self._round) + + def _copy_landed(self, lease: _StagingLease) -> bool: + """A read lease's copy into its slots has completed.""" + return self._landed(lease._copy) + + def _recyclable(self, lease: _StagingLease) -> bool: + """The lease's slots may return: no backend access possible and no copy on them pending. + The copy is asked last, so a lease still lent costs no event query.""" + # A failed lease was never ready, so no backend has seen its slots. + if lease.failure is None: + if not lease._released: + return False + if not (lease._kind == "read" or lease._marked or not lease._seen_ready): + return False + # TODO: every lender call asks each released lease's pending copy again, so N calls while P + # copies pend cost N*P event queries; copies on one stream complete in order, so asking + # from the oldest and stopping at the first pending one would cost about one per call. + return self._landed(lease._copy) + + def _still_lent(self, kv, lease: _StagingLease) -> List[np.ndarray]: + """Per run, the rows whose block is still in the window of the same active cache and still + locks the GPU page lent; a page only held may sit on another tier under the same number.""" + rows = lease._rows + manager = self._manager_ref() + state = _manager.cache_state(kv) if kv is not None and kv is lease._kv else None + masks = [] + for lg, ordinals, slots in zip(rows.layer_groups, rows.ordinals, rows.device_slots): + same = np.zeros(len(ordinals), dtype=bool) + if state is not None and state.active: + pages = _manager.locked_pages(kv, lg) + inside = ordinals < len(pages) + same[inside] = pages[ordinals[inside]] == slots[inside] + stale_beg, stale_end = _stale(manager, self._layout, lg, state.history) + same &= (ordinals < stale_beg) | (ordinals >= stale_end) + masks.append(same) + return masks + + def _scratch_reuse_on(self, kv) -> bool: + """A windowed write target whose window blocks may sit in scratch slots, which a fetch into + them would not survive.""" + return self._any_window and _manager.scratch_reuse(kv) + + def _report_settled(self, fetch: _Fetch) -> bool: + """The fetch's arrived rows are marked and their copy into the request's pages is done.""" + return fetch.delivered and self._landed(fetch.copy) + + def _fail_waiting(self, request_id: int, reason: str) -> None: + """Fail the request's leases still waiting for slots.""" + for lease in [lease for lease in self._line if lease._request_id == request_id]: + self._fail(lease, reason) + + def _abandon(self, lease: _StagingLease) -> None: + """Drop the write lease's fetch record: it delivered nothing, and readiness counts what + earlier fetches into the cache delivered.""" + if lease._fetch is not None and self._fetches.get(lease._request_id) is lease._fetch: + del self._fetches[lease._request_id] + + def _names(self, layer_group: int, keys: np.ndarray) -> np.ndarray: + """The rows' names: ``uint8 (n, 54)`` for ``keys`` ``uint8 (n, 32)``.""" + return self._identity.names(layer_group, keys) + + def _memory_in_use(self) -> bool: + """A lease or a hold is unreleased, or a slot is lost to a failed copy: the memory stays.""" + return bool(self._unreleased) or bool(self._holds) or bool(self._quarantined) + + def _keep_index_buffer(self, manager) -> None: + """Keep the manager's page-index buffer until its shutdown: a cache that a lease or a record + holds when the manager goes without one writes its page indices there as it closes.""" + self._index_buffer = _manager.index_buffer(manager) + _KEPT.append(self._index_buffer) + + def _let_go_index_buffer(self) -> None: + """At the manager's shutdown: it closes every cache after this hook, while it still holds the + page-index buffer, so no cache writes there later.""" + _let_go(self._index_buffer) + self._index_buffer = None + + def _end_hold(self, hold: _PartsHold) -> None: + """The hold's holder let go; ``KeyError`` for a hold not open, which its guard prevents.""" + self._holds.remove(hold) + + def _check_fits(self, counts: Mapping[int, int]) -> None: + """``ValueError`` if a lease needing ``counts`` rows per pool group can never be granted.""" + try: + self._slots.check(counts) + except ValueError as error: + raise ValueError( + f"the range needs more staging slots than there are ({error}); ranges of at most " + "fetch_tokens tokens always fit" + ) from None + + def _free_slots(self, group: int) -> int: + """Free slots of a pool group; for tests.""" + return self._slots.free_slots(group) + + def _open_count(self) -> int: + """Leases not yet released; for tests.""" + return len(self._unreleased) + + # Lease records, slots and grants. + + def _live(self): + """The manager while the lender serves it; ``None`` once it shut down or is gone.""" + if self._closed: + return None + return self._manager_ref() + + def _whole_blocks(self, start: int, end: int) -> Tuple[int, int]: + start, end = int(start), int(end) + tpb = int(self._layout.tokens_per_block) + if start < 0 or end < 0 or start > end: + raise ValueError(f"bad token range [{start}, {end})") + if start % tpb or end % tpb: + raise ValueError( + f"a staging lease covers whole blocks of {tpb} tokens, got [{start}, {end})" + ) + return start, end + + def _counts(self, ordinals: Sequence[np.ndarray]) -> Dict[int, int]: + counts: Dict[int, int] = {} + for lg, lg_ordinals in enumerate(ordinals): + g = int(self._layout.pool_group_of[lg]) + counts[g] = counts.get(g, 0) + len(lg_ordinals) + return counts + + def _open(self, lease: _StagingLease, counts: Mapping[int, int]) -> None: + """Record a new lease and grant it now when no lease waits and its slots are free.""" + self._unreleased.add(lease) + if lease._rows.num_rows == 0: + # Nothing to stage: ready at the first poll, without slots or a place in line. + self._grant(lease, None) + return + lease._ticket = self._slots.ask(counts) + self._line.append(lease) + self._grant_waiting(fresh=lease) + + def _grant_waiting(self, fresh: Optional[_StagingLease] = None) -> None: + """Grant the leases in line, strictly in order; ``fresh`` was looked up in this call.""" + while self._line: + head = self._line[0] + runs = self._slots.take(head._ticket) + if runs is None: + return + self._line.popleft() + head._ticket = None + try: + self._grant(head, runs, recheck=head is not fresh) + except Exception as error: + # One broken grant fails only its own lease: it neither stalls the line nor raises + # out of another lease's call. Slots a copy may have been queued on stay unused. + if any(lease is head for lease in self._holding): + self._quarantine(head) + else: + self._slots.give(runs) + self._fail(head, f"granting staging slots failed: {error!r}") + logger.warning(f"KV cache lender: {head.failure}") + + def _grant(self, lease: _StagingLease, runs: Optional[Runs], recheck: bool = False) -> None: + """Give ``lease`` its slots and view; a read then queues its copy into them.""" + if recheck and lease._kind == "read": + problem = self._source_changed(lease) + if problem is not None: + self._slots.give(runs) + self._fail(lease, problem) + return + self._assign_staging(lease._rows, runs) + lease._view = self._view(lease._rows, lease._keys) + lease._granted = True + if runs is None: + return + lease._runs = runs + self._holding.append(lease) + if lease._kind != "read": + return + copy, error = self._memcpy(self._segments(lease._rows), to_staging=True) + if copy is None: + self._quarantine(lease) + else: + lease._copy = copy + if error is not None: + self._fail(lease, f"the copy into staging failed: {error}") + + def _source_changed(self, lease: _StagingLease) -> Optional[str]: + """Why a read granted after waiting cannot copy the pages it looked up any more, if so.""" + manager = self._manager_ref() + kv = _manager.kv_of(manager, lease._request_id) + if kv is not lease._kv: + return "the request's cache was freed while the lease waited" + state = _manager.cache_state(kv) + if not state.active: + return "the request's cache was suspended while the lease waited" + rows = lease._rows + for lg, ordinals, slots in zip(rows.layer_groups, rows.ordinals, rows.device_slots): + pages = _manager.locked_pages(kv, lg) + if np.any(ordinals >= len(pages)) or np.any(pages[ordinals] != slots): + return f"layer group {lg}: pages changed while the lease waited" + stale_beg, stale_end = _stale(manager, self._layout, lg, state.history) + if np.any((ordinals >= stale_beg) & (ordinals < stale_end)): + return f"layer group {lg}: blocks left the window while the lease waited" + return None + + def _fail(self, lease: _StagingLease, reason: str) -> None: + """Fail a lease never seen ready: it leaves the line, and a write abandons its fetch.""" + if lease.failure is not None: + return + lease._set_failure(reason) + if lease._ticket is not None: + self._slots.cancel(lease._ticket) + lease._ticket = None + self._line.remove(lease) + if lease._kind == "write": + self._abandon(lease) + + def _quarantine(self, lease: _StagingLease) -> None: + """Keep the lease's slots from reuse for good: a copy on them may still be running.""" + self._holding = [held for held in self._holding if held is not lease] + # TODO: reclaim quarantined slots, or alert when they erode staging capacity. + if lease._runs is not None: + self._quarantined.append(lease._runs) + logger.warning( + f"KV cache lender: a failed copy took staging slots of request {lease._request_id}" + ) + + def _recycle(self) -> None: + """Return the slots of every lease ``_recyclable`` allows.""" + holding = [] + for lease in self._holding: + if self._recyclable(lease): + self._slots.give(lease._runs) + else: + holding.append(lease) + self._holding = holding + + def _on_release(self, lease: _StagingLease) -> None: + """The lease's holder let go. The record changes first and progress runs last, so an + unexpected error leaves the fetch abandoned rather than unsettled.""" + self._unreleased.discard(lease) + if self._live() is None: + return + if lease._ticket is not None: + self._fail(lease, "released while waiting for staging slots") + elif lease._kind == "write" and not lease._seen_ready and lease.failure is None: + # Released before anyone saw it ready: no backend wrote, and no marks are due. + lease._doomed = None + self._abandon(lease) + self._progress() + + def _apply_marks(self, lease: _StagingLease, masks: List[np.ndarray]) -> None: + """Copy the marked rows whose page is still the one lent into the request's pages. The + fetch stays abandoned until that copy is queued without error; then its rows add to the + cache's deliveries.""" + fetch = lease._fetch + current = fetch is not None and self._fetches.get(lease._request_id) is fetch + if current: + del self._fetches[lease._request_id] + kv = _manager.kv_of(self._manager_ref(), lease._request_id) + lent = self._still_lent(kv, lease) + copy = [mask & still for mask, still in zip(masks, lent)] + error = None + if any(c.any() for c in copy): + queued, error = self._memcpy(self._segments(lease._rows, copy), to_staging=False) + if queued is None: + self._quarantine(lease) + else: + lease._copy = queued + if current and error is None: + fetch.copy = lease._copy + fetch.delivered = True + self._fetches[lease._request_id] = fetch + self._deliver(lease._request_id, fetch, lease._rows, copy) + self._progress() + + def _deliver( + self, request_id: int, fetch: _Fetch, rows: _Rows, copied: List[np.ndarray] + ) -> None: + """Add a fetch's copied rows to what earlier fetches into the same cache delivered, so a + fetch split into consecutive leases counts as one.""" + delivered = self._delivered.get(request_id) + if delivered is None or delivered.kv is not fetch.kv: + empty = [np.zeros(0, dtype=bool) for _ in range(self._layout.num_layer_groups)] + delivered = _Delivered(fetch.kv, fetch.start, empty) + self._delivered[request_id] = delivered + delivered.origin = min(delivered.origin, fetch.start) + for lg, ordinals, mask in zip(rows.layer_groups, rows.ordinals, copied): + got = ordinals[mask] + if not len(got): + continue + blocks = delivered.blocks[lg] + if int(got.max()) >= len(blocks): + blocks = np.concatenate([blocks, np.zeros(int(got.max()) + 1 - len(blocks), bool)]) + blocks[got] = True + delivered.blocks[lg] = blocks + delivered.usable = None + + def _void_past(self, delivered: _Delivered, kv) -> None: + """Forget delivered rows at or past the cache's block count: a shrink freed their pages, + and a regrow brings pages without their contents.""" + kept = _manager.num_blocks(kv) + for blocks in delivered.blocks: + if blocks[kept:].any(): + blocks[kept:] = False + delivered.usable = None + + def _rows(self, manager, kv, start: int, end: int) -> _Rows: + """The blocks of ``[start, end)`` a history of ``end`` reads, per layer group, with their + device slots (-1 where a block has no page).""" + layout = self._layout + tpb = int(layout.tokens_per_block) + rows = _Rows([], [], []) + for lg in range(layout.num_layer_groups): + ordinals = _needed_ordinals(manager, layout, lg, start // tpb, end // tpb, end) + pages = _manager.pages(kv, lg) + slots = np.full(len(ordinals), -1, dtype=np.int64) + inside = ordinals < len(pages) + slots[inside] = pages[ordinals[inside]] + rows.layer_groups.append(lg) + rows.ordinals.append(ordinals) + rows.device_slots.append(slots) + return rows + + def _lendable(self, manager, history: int, rows: _Rows) -> _Rows: + """``rows`` without blocks that have no page or that the request's own window has passed + (their page may hold something else).""" + out = _Rows([], [], []) + for lg, ordinals, slots in zip(rows.layer_groups, rows.ordinals, rows.device_slots): + stale_beg, stale_end = _stale(manager, self._layout, lg, history) + ok = (slots >= 0) & ~((ordinals >= stale_beg) & (ordinals < stale_end)) + out.layer_groups.append(lg) + out.ordinals.append(ordinals[ok]) + out.device_slots.append(slots[ok]) + return out + + def _keys_for( + self, manager, request: LlmRequest, kv, ordinals: Sequence[np.ndarray] + ) -> List[np.ndarray]: + """Per layer group, the reuse keys of its rows' blocks, ``uint8 (n, 32)``.""" + # TODO: every lease hashes the request's whole prefix again from block 0, so a lease late in + # a long prompt costs time in proportion to the prompt; cache the key chain per request. + top = max((int(o.max()) + 1 for o in ordinals if len(o)), default=0) + keys = _manager.block_keys(manager, request, kv, top) + return [_key_column([keys[int(o)] for o in lg_ordinals]) for lg_ordinals in ordinals] + + def _assign_staging(self, rows: _Rows, runs: Optional[Runs]) -> None: + """Each pool group's run goes to its layer groups in order, so the rows of one pool group + occupy consecutive slots.""" + cursor: Dict[int, int] = {} + rows.staging_slots = [] + for lg, ordinals in zip(rows.layer_groups, rows.ordinals): + g = int(self._layout.pool_group_of[lg]) + start = runs.runs.get(g, (0, 0))[0] if runs is not None else 0 + offset = cursor.get(g, 0) + rows.staging_slots.append( + np.arange(start + offset, start + offset + len(ordinals), dtype=np.int64) + ) + cursor[g] = offset + len(ordinals) + + def _view(self, rows: _Rows, keys: List[np.ndarray]) -> RegionView: + runs = [] + for lg, ordinals, lg_keys, slots in zip( + rows.layer_groups, rows.ordinals, keys, rows.staging_slots + ): + index = self._part_of_group[int(self._layout.pool_group_of[lg])] + part = self._parts[index] + runs.append( + GroupRun( + lg, + ordinals, + names=self._names(lg, lg_keys), + addresses=part.address + slots * part.slot_bytes, + part=index, + ) + ) + return RegionView(tuple(runs)) + + def _group_rows(self, rows: _Rows, mask: Optional[List[np.ndarray]] = None): + """Rows by device pool group: ``(g, device slots, staging slots)``, filtered by ``mask`` + (one boolean array per layer group; ``None``: every row).""" + by_group: Dict[int, List[int]] = {} + for i, lg in enumerate(rows.layer_groups): + by_group.setdefault(int(self._layout.pool_group_of[lg]), []).append(i) + for g, members in by_group.items(): + dev = np.concatenate([rows.device_slots[i] for i in members]) + stg = np.concatenate([rows.staging_slots[i] for i in members]) + if mask is not None: + keep = np.concatenate([np.asarray(mask[i], dtype=bool) for i in members]) + dev, stg = dev[keep], stg[keep] + if len(dev): + yield g, dev, stg + + def _segments(self, rows: _Rows, mask: Optional[List[np.ndarray]] = None) -> List[List[int]]: + """``[staging address, device address, bytes]`` for every pool of every row: a staging slot + holds the device pools' slots back to back in pool order. Segments that continue each other + on both sides are merged.""" + out: List[List[int]] = [] + for g, dev, stg in self._group_rows(rows, mask): + part = self._parts[self._part_of_group[g]] + offset = 0 + for pool in self._layout.device_pools[g]: + width = int(pool.slot_bytes) + for d, t in zip(dev.tolist(), stg.tolist()): + host = part.address + t * part.slot_bytes + offset + device = int(pool.base) + d * width + last = out[-1] if out else None + if last and last[0] + last[2] == host and last[1] + last[2] == device: + last[2] += width + else: + out.append([host, device, width]) + offset += width + return out + + # TODO: move the copies to a side stream ordered by events once every path that releases a page + # waits for them, and issue fewer copy calls, batched or merged. + def _memcpy( + self, segments: Sequence[Sequence[int]], to_staging: bool + ) -> Tuple[Optional[_Copy], Optional[str]]: + """Queue async copies on the manager's execution stream, then record an event covering + every copy queued, even after one failed; ``(None, error)`` if it could record none.""" + # The execution stream orders a copy after the forward passes that wrote its pages and + # before any later owner of those pages writes them: serial with the forward on the GPU. + # With page-locked staging the CPU does not wait for it; pageable staging may hold the call. + stream = _manager.stream(self._manager_ref()) + handle = drv.CUstream(stream.cuda_stream) + error = None + for host, device, nbytes in segments: + dst, src = (host, device) if to_staging else (device, host) + (result,) = drv.cuMemcpyAsync( + drv.CUdeviceptr(dst), drv.CUdeviceptr(src), nbytes, handle + ) + if result != drv.CUresult.CUDA_SUCCESS: + error = f"cuMemcpyAsync of {nbytes} bytes failed: {result}" + break + try: + event = torch.cuda.Event() + event.record(stream) + except RuntimeError as record_error: + error = error or f"recording the copy's event failed: {record_error}" + logger.warning(f"KV cache lender: {error}") + return None, error + if error is not None: + logger.warning(f"KV cache lender: {error}") + return _Copy(event), error + + def _floor(self, manager, history: int, origin: int) -> int: + """The lowest start that needs no restart: the history when a window has released blocks + at it (they have no pages), else the smaller of the first fetch start and the history.""" + for lg in range(self._layout.num_layer_groups): + stale_beg, stale_end = _stale(manager, self._layout, lg, history) + if stale_end > stale_beg: + return history + return min(origin, history) + + def _usable_until(self, manager, delivered: _Delivered, committed: int) -> int: + """The largest start ``P >= committed`` where every layer group has what it reads among the + committed blocks and delivered rows: full attention every block below ``P``, a window its + sinks and in-window blocks. Non-monotonic in ``P`` under windows, so each is checked.""" + tpb = int(self._layout.tokens_per_block) + if delivered.origin > committed: + return committed # the first fetch left a gap after the committed tokens + first = committed // tpb # the blocks below are committed + last = max((len(blocks) for blocks in delivered.blocks), default=0) + if last <= first: + return committed + ok = np.ones(last - first, dtype=bool) # ok[i]: start at (first + 1 + i) * tpb + for lg, blocks in enumerate(delivered.blocks): + have = np.zeros(last - first, dtype=bool) + mine = blocks[first:last] + have[: len(mine)] = mine + missing = np.concatenate([[0], np.cumsum(~have)]) # missing in [first, first + j) + + def gap(a: int, b: int) -> bool: + a, b = max(a, first), min(b, last) + return b > a and missing[b - first] - missing[a - first] > 0 + + for i, end_block in enumerate(range(first + 1, last + 1)): + if not ok[i]: + continue + stale_beg, stale_end = _stale(manager, self._layout, lg, end_block * tpb) + if stale_end > stale_beg: + bad = gap(first, min(stale_beg, end_block)) or gap(stale_end, end_block) + else: + bad = gap(first, end_block) + ok[i] = not bad + good = np.nonzero(ok)[0] + return committed if not len(good) else (first + 1 + int(good[-1])) * tpb + + +# TODO: an in-place lease gets none of the staging guarantees: nothing guards its pages against +# suspend, shrink, window advance or a pool rebalance while a loan is open, it is ready before the +# manager's stream finishes with them, and mark_arrived feeds no readiness; the caller stands in. +class InPlace: + """``InPlaceLender`` over one manager, which it references weakly. A lease holds a loan on the + request's cache; the request's free keeps a lent cache open until its last loan ends.""" + + def __init__(self, manager: weakref.ref, layout: ManagerLayout) -> None: + self._manager_ref = manager + self._layout = layout + # Open loans per cache, holding the cache strongly so that dropping a lease never lets a + # collector close it on another thread. + self._loans: Dict[object, int] = {} + self._freed: Dict[object, Callable[[], None]] = {} # lent caches the manager freed + self._kept: Optional[FrozenSet[object]] = None # set by the manager's shutdown + # Kept while a loan is open, and until exit once the manager is gone. + self._index_buffer: Optional[object] = None + + def lend_read(self, request: LlmRequest, start: int, end: int) -> _InPlaceLease: + """See ``InPlaceLender.lend_read``.""" + return self._lend(request, start, end, "read") + + def lend_write(self, request: LlmRequest, start: int, end: int) -> _InPlaceLease: + """See ``InPlaceLender.lend_write``.""" + return self._lend(request, start, end, "write") + + def _on_free(self, request_id: int, kv_cache, after_close: Callable[[], None]) -> bool: + """Manager hook: ``True`` if ``kv_cache`` is on loan, which the release ending its last loan + then closes, running ``after_close`` in that call. Never raises.""" + try: + if kv_cache not in self._loans: + return False + self._freed[kv_cache] = after_close + return True + except Exception: + # A hook never raises into the manager's free. It keeps the cache open while any loan + # is, since closing a lent cache would hand its pages to another request. + logger.error(f"KV cache lender: freeing request {request_id}: {traceback.format_exc()}") + return bool(self._loans) + + def _on_shrink(self, request_id: int, kv_cache) -> None: + """Manager hook after an in-place shrink: nothing to do, since the caller keeps a lent + cache from shrinking.""" + + def _on_shutdown(self, impl) -> FrozenSet[object]: + """Manager hook: the caches still on loan, kept with ``impl`` until the process exits; the + same set on every later call.""" + if self._kept is not None: + return self._kept + self._kept = frozenset(self._loans) + if self._kept: + try: + # A device pool cannot be freed in part, so the pools stay with the lent caches. + _KEPT.extend([impl, *self._kept]) + logger.warning( + f"KV cache lender: keeping {len(self._kept)} lent caches and their pools " + "until exit" + ) + except Exception: + # A hook never raises into the manager's shutdown, which still leaves the lent + # caches open and the pools unfreed. + logger.error(f"KV cache lender: shutting down: {traceback.format_exc()}") + return self._kept + + def _end_loan(self, kv_cache) -> None: + """End one loan on ``kv_cache``; the last loan on a freed cache closes it in this call.""" + if self._kept is not None: + return # every cache on loan at the manager's shutdown is kept until exit + left = self._loans.get(kv_cache, 0) - 1 + if left > 0: + self._loans[kv_cache] = left + return + self._loans.pop(kv_cache, None) + after_close = self._freed.pop(kv_cache, None) + if after_close is not None: + _manager.close_cache(kv_cache) + after_close() + if not self._loans: + self._let_go_index_buffer() + + def _keep_index_buffer(self, manager) -> None: + """Keep the manager's page-index buffer while a loan is open: a lent cache the manager does + not detach writes its page indices there as it closes, even after the manager is gone.""" + self._index_buffer = _manager.index_buffer(manager) + _KEPT.append(self._index_buffer) + + def _let_go_index_buffer(self) -> None: + """The last loan ended: let the buffer go to its manager. With the manager gone, a cache + this lender held may still close after this call, so the buffer stays until exit.""" + if self._manager_ref() is None: + return + _let_go(self._index_buffer) + self._index_buffer = None + + def _lend(self, request: LlmRequest, start: int, end: int, kind: str) -> _InPlaceLease: + start, end = int(start), int(end) + if start < 0 or end < 0 or start > end: + raise ValueError(f"bad token range [{start}, {end})") + request_id = int(request.py_request_id) + manager = self._manager_ref() + if self._kept is not None or manager is None: + return _InPlaceLease._failed(self, kind, _SHUT_DOWN) + kv = _manager.kv_of(manager, request_id) + if kv is None: + return _InPlaceLease._failed(self, kind, _no_cache(request_id)) + state = _manager.cache_state(kv) + if not state.active: + return _InPlaceLease._failed(self, kind, _SUSPENDED) + runs = [] + for lg in range(self._layout.num_layer_groups): + ordinals, slots = self._device_slots(manager, kv, lg, start, end) + paged = slots >= 0 + if kind == "write" and not paged.all(): + missing = ordinals[~paged].tolist() + return _InPlaceLease._failed( + self, kind, f"layer group {lg}: blocks {missing[:8]} have no page" + ) + # A read leaves blocks without a page out. + runs.append(GroupRun(lg, ordinals[paged])) + # The loan opens here: from now on the request's free keeps the cache open. + if not self._loans: + self._keep_index_buffer(manager) + self._loans[kv] = self._loans.get(kv, 0) + 1 + return _InPlaceLease(self, kind, RegionView(tuple(runs)), kv) + + def _device_slots( + self, manager, kv, lg: int, start: int, end: int + ) -> Tuple[np.ndarray, np.ndarray]: + """The blocks of ``lg`` that ``[start, end)`` touches, a partial last one included, that a + history of ``end`` reads, and their device slots (-1 where a block has no locked page).""" + layout = self._layout + tpb = int(layout.tokens_per_block) + first, last = start // tpb, -(-end // tpb) + windowed = layout.windows[lg] is not None + # Only pages the cache locks: a window block behind its history keeps at most a held page, + # which a lower cache tier may take at any time. + pages = _manager.locked_pages(kv, lg) + if windowed: + ordinals = _needed_ordinals(manager, layout, lg, first, last, end) + slots = np.full(len(ordinals), -1, dtype=np.int64) + inside = ordinals < len(pages) + slots[inside] = pages[ordinals[inside]] + else: + ordinals = np.arange(first, last, dtype=np.int64) + slots = np.full(len(ordinals), -1, dtype=np.int64) + paged = max(0, min(last, len(pages)) - first) + slots[:paged] = pages[first : first + paged] + return ordinals, slots + + +class _StagingLease: + """A staging lease; holds its lender weakly and has no finalizer. Backends may read its view's + arrays on their own threads until release, and never call it.""" + + def __init__( + self, + lender: Staging, + kind: str, + request_id: int, + kv=None, + rows: Optional[_Rows] = None, + keys: Optional[List[np.ndarray]] = None, + fetch: Optional[_Fetch] = None, + ) -> None: + self._lender = weakref.ref(lender) + self._kind = kind + self._request_id = request_id + self._kv = kv + self._rows = rows + self._keys = keys + self._fetch = fetch + self._doomed: Optional[str] = None # why it fails at its first poll + self._ticket: Optional[int] = None # its place in line while it waits for slots + self._runs: Optional[Runs] = None + self._view: Optional[RegionView] = None + self._copy: Optional[_Copy] = None + self._granted = False + self._seen_ready = False + self._released = False + self._marked = False + self._failure: Optional[str] = None + + @classmethod + def _failed(cls, lender: Staging, kind: str, request_id: int, reason: str) -> _StagingLease: + """A lease failed at the call; open until released, like any other.""" + lease = cls(lender, kind, request_id) + lease._set_failure(reason) + if not lender._closed: + lender._unreleased.add(lease) + return lease + + def _set_failure(self, reason: str) -> None: + self._failure = reason + self._doomed = None + + def poll(self) -> Optional[RegionView]: + """See ``Lease.poll``.""" + if self._released: + raise RuntimeError("poll after release") + lender = self._lender() + if lender is None or lender._live() is None: + return self._ended_poll() + lender._progress() + if self._failure is not None: + return None + if self._doomed is not None: + lender._fail(self, self._doomed) + return None + if not self._granted: + return None + if self._kind == "read" and not lender._copy_landed(self): + return None + self._seen_ready = True + return self._view + + @property + def failure(self) -> Optional[str]: + """See ``Lease.failure``.""" + return self._failure + + def mark_arrived(self, masks: Sequence[np.ndarray]) -> None: + """See ``Lease.mark_arrived``.""" + if self._kind != "write": + raise RuntimeError("mark_arrived is for write leases") + if not self._seen_ready: + raise RuntimeError("mark_arrived before poll() returned the view") + if self._marked: + raise RuntimeError("mark_arrived called twice") + masks = _checked_masks(self._view, masks) + self._marked = True + lender = self._lender() + if lender is None or lender._live() is None: + return + lender._apply_marks(self, masks) + + def release(self) -> None: + """See ``Lease.release``.""" + if self._released: + return + self._released = True + lender = self._lender() + if lender is not None: + lender._on_release(self) + + def _ended_poll(self) -> Optional[RegionView]: + """``poll`` after the lender stopped serving: no grant or copy, just what already landed.""" + if self._failure is not None or not self._granted: + return None + if self._kind == "read" and self._copy is not None and not self._copy.done(): + return None + self._seen_ready = True + return self._view + + +class _InPlaceLease: + """An in-place lease: the loan on the request's cache, held until release. Holds its lender + weakly and has no finalizer.""" + + def __init__( + self, + lender: InPlace, + kind: str, + view: Optional[RegionView], + kv_cache=None, + failure: Optional[str] = None, + ) -> None: + self._lender = weakref.ref(lender) + self._kind = kind + self._view = view + self._cache = kv_cache # the loan, until release + self._failure = failure + self._seen_ready = False + self._released = False + self._marked = False + + @classmethod + def _failed(cls, lender: InPlace, kind: str, reason: str) -> _InPlaceLease: + """A lease failed at the call; it holds no loan.""" + return cls(lender, kind, None, failure=reason) + + def poll(self) -> Optional[RegionView]: + """See ``Lease.poll``.""" + if self._released: + raise RuntimeError("poll after release") + if self._failure is not None: + return None + # The request's own pages: ready at once; the caller has let the stream's work on them end. + self._seen_ready = True + return self._view + + @property + def failure(self) -> Optional[str]: + """See ``Lease.failure``.""" + return self._failure + + def mark_arrived(self, masks: Sequence[np.ndarray]) -> None: + """See ``Lease.mark_arrived``.""" + if self._kind != "write": + raise RuntimeError("mark_arrived is for write leases") + if not self._seen_ready: + raise RuntimeError("mark_arrived before poll() returned the view") + if self._marked: + raise RuntimeError("mark_arrived called twice") + # The rows are in the request's pages already: only the shapes are checked. + _checked_masks(self._view, masks) + self._marked = True + + def release(self) -> None: + """See ``Lease.release``.""" + if self._released: + return + self._released = True + cache, self._cache = self._cache, None + lender = self._lender() + if cache is not None and lender is not None: + lender._end_loan(cache) + + +class _PartsHold: + """A hold on the staging memory; holds its lender weakly and has no finalizer. Without a lender + it is inert.""" + + def __init__(self, lender: Optional[Staging]) -> None: + self._lender = weakref.ref(lender) if lender is not None else None + self._released = False + + def release(self) -> None: + """See ``PartsHold.release``.""" + if self._released: + return + self._released = True + lender = self._lender() if self._lender is not None else None + if lender is not None: + lender._end_hold(self) diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_manager.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_manager.py new file mode 100644 index 000000000000..969006e98caa --- /dev/null +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_manager.py @@ -0,0 +1,191 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Every read or write of manager and runtime internals the lender makes, one function each, so +moving the lender into the manager or changing a member touches this file alone. Imports of the +manager and the runtime stay inside the functions.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Dict, List, NamedTuple, Optional, Tuple + +import numpy as np + +if TYPE_CHECKING: + import torch + + from ...llm_request import LlmRequest + from ..kv_cache_manager_v2 import KVCacheManagerV2 + + +class CacheState(NamedTuple): + """A request's runtime cache as the lender reads it: tokens committed, the history length its + sliding windows keep, and whether it is active (not suspended).""" + + committed: int + history: int + active: bool + + +def require_v2(manager) -> None: + """Raise ``TypeError`` unless ``manager`` is a ``KVCacheManagerV2`` (subclasses included).""" + from ..kv_cache_manager_v2 import KVCacheManagerV2 + + if not isinstance(manager, KVCacheManagerV2): + raise TypeError( + f"the KV cache lender needs a KVCacheManagerV2, got {type(manager).__name__}" + ) + + +def commits_blocks(manager: KVCacheManagerV2) -> bool: + """Whether the manager commits blocks to its prefix-reuse tree: block reuse on and, for a draft + manager, joint reuse with its target.""" + return bool(manager.enable_block_reuse) and bool(manager._can_publish_block_reuse) + + +def attached(manager: KVCacheManagerV2) -> Optional[object]: + """The lender attached to ``manager``, or ``None``.""" + return getattr(manager, "_sharing", None) + + +def install(manager: KVCacheManagerV2, lender: object) -> None: + """Attach ``lender`` for the manager's life: its free, in-place shrinks and shutdown then call + the lender's ``_on_free``, ``_on_shrink`` and ``_on_shutdown``.""" + manager._sharing = lender + + +def index_buffer(manager: KVCacheManagerV2) -> torch.Tensor: + """The host buffer the manager's caches write their page indices into, through a raw pointer + each keeps until it is detached or closed, whether or not the manager still exists.""" + return manager.host_kv_cache_block_offsets + + +def virtual_layers(manager: KVCacheManagerV2) -> Optional[Tuple[Dict[int, Tuple[int, int]], int]]: + """Virtual layers (DeepSeek-V4): internal layer id -> (model layer, attention type value), and + the number of attention types the enum defines. ``None`` for a manager without them.""" + virtual = getattr(manager, "_layer_attn_to_layer_id", None) + if not virtual: + return None + inverse: Dict[int, Tuple[int, int]] = {} + for (model_layer, attn_type), layer_id in virtual.items(): + inverse[int(layer_id)] = (int(model_layer), int(attn_type.value)) + # Every member of the enum counts, so stages holding different attention types agree. + attn_type_class = type(next(iter(virtual))[1]) + return inverse, max(int(member.value) for member in attn_type_class) + 1 + + +def stream(manager: KVCacheManagerV2) -> torch.cuda.Stream: + """The manager's execution stream: a copy queued on it runs after the forward passes that wrote + the pages and before the pages' next writer.""" + return manager._stream + + +def kv_of(manager: KVCacheManagerV2, request_id: int) -> Optional[object]: + """The request's runtime cache, or ``None`` when it has none (never had one, freed, shut down). + Caches are compared by identity to see a free or a replaced cache.""" + return manager.kv_cache_map.get(request_id) + + +def cache_state(kv) -> CacheState: + """Committed tokens, history length and activity of runtime cache ``kv``, as plain values.""" + return CacheState(int(kv.num_committed_tokens), int(kv.history_length), bool(kv.is_active)) + + +def scratch_reuse(kv) -> bool: + """Whether runtime cache ``kv`` has SWA scratch reuse on: its window blocks may sit in scratch + slots, which the next chunk overwrites.""" + return bool(kv.enable_swa_scratch_reuse) + + +def pages(kv, layer_group: int) -> np.ndarray: + """``int64`` page (pool-group slot) of each block ordinal of ``layer_group``, -1 where the block + has no page.""" + return np.fromiter( + kv.get_aggregated_page_indices(layer_group, valid_only=False), dtype=np.int64 + ) + + +def locked_pages(kv, layer_group: int) -> np.ndarray: + """``int64`` page of each block of an active cache from its locked base page indices, -1 where + the block has none: a window block behind the history only holds its page, which may be on + another tier.""" + return np.array(kv.get_base_page_indices(layer_group)[: kv.num_blocks], dtype=np.int64) + + +def num_blocks(kv) -> int: + """Block ordinals of runtime cache ``kv``: a shrink drops those past its capacity, with their + pages.""" + return int(kv.num_blocks) + + +def stale_blocks(manager: KVCacheManagerV2, layer_group: int, history: int) -> Tuple[int, int]: + """Block ordinals ``[beg, end)`` behind a windowed layer group's window at ``history``, sinks + excepted. Called only for layer groups with a window.""" + beg, end = manager._stale_block_range(layer_group, history) + return int(beg), int(end) + + +def block_keys(manager: KVCacheManagerV2, request: LlmRequest, kv, num_blocks: int) -> List[bytes]: + """The 32-byte reuse keys of the request's first ``num_blocks`` whole blocks, over the tokens + and reuse scope the manager commits with. ``ValueError`` if the request has fewer.""" + if num_blocks <= 0: + return [] + from tensorrt_llm.runtime.kv_cache_manager_v2 import sequence_to_blockchain_keys + + tpb = int(manager.tokens_per_block) + source = manager._reuse_token_source(request) + need = num_blocks * tpb + if len(source) < need: + raise ValueError( + f"request {request.py_request_id} has {len(source) // tpb} whole blocks of tokens, " + f"{num_blocks} asked" + ) + # Multimodal digests take the place of placeholder tokens, as when the manager commits. + tokens = manager._augment_tokens_for_block_reuse(source, request, 0, need) + if isinstance(tokens, np.ndarray): + tokens = tokens.tolist() + keys: List[bytes] = [] + for i, (block_tokens, key) in enumerate( + sequence_to_blockchain_keys(tpb, kv.reuse_scope, tokens) + ): + if i == 0: + continue # the root: the reuse scope's own key + if len(block_tokens) < tpb: + break + keys.append(bytes(key)) + if len(keys) == num_blocks: + break + if len(keys) < num_blocks: + raise ValueError( + f"request {request.py_request_id} has {len(keys)} whole blocks of tokens, " + f"{num_blocks} asked" + ) + return keys + + +def grow(manager: KVCacheManagerV2, request: LlmRequest, kv, position: int, end: int) -> bool: + """Cover ``[0, end)`` and move the history to ``position`` with the manager's own resize, then + run its fresh-page fill (a diagnostic, off unless set), as every resize of the manager does. + ``False``, the cache unchanged, when pages run out.""" + if not manager._resize_for_connector_prefix(request, kv, position, end): + return False + # The fill marks the new pages before a fetch lands in them; a later resize then sees them as + # the request's own and leaves the fetched blocks alone. + manager._fill_fresh_kv_pages(request.py_request_id) + return True + + +def close_cache(kv) -> None: + """Close runtime cache ``kv``: its pages go back to the pools.""" + kv.close() diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_slots.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_slots.py new file mode 100644 index 000000000000..d27cca74c6e0 --- /dev/null +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_slots.py @@ -0,0 +1,233 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Staging sizing and slot bookkeeping, without memory. A lease takes all its slots at once, one +contiguous run per pool group, or none; leases waiting in line are served strictly in order.""" + +from __future__ import annotations + +import bisect +from collections import OrderedDict +from dataclasses import dataclass +from typing import TYPE_CHECKING, Dict, List, Mapping, Optional, Tuple + +import numpy as np + +from ._types import StagingOptions, _positive + +if TYPE_CHECKING: + from ._layout import ManagerLayout + + +def fetch_rows(layout: ManagerLayout, fetch_tokens: int) -> Dict[int, int]: + """Rows per device pool group that any whole-block range of at most ``fetch_tokens`` tokens can + need: ``ceil(fetch_tokens / tpb)`` per layer group, with a window at most its sink blocks plus + ``ceil(window / tpb)``. ``ValueError`` if ``fetch_tokens < 1``, ``TypeError`` if not an int.""" + tpb = int(layout.tokens_per_block) + target = -(-_positive("fetch_tokens", fetch_tokens) // tpb) + rows: Dict[int, int] = {} + for lg, group in enumerate(layout.pool_group_of): + blocks = target + window = layout.windows[lg] + # A window of W tokens reads at most ceil(W / tpb) blocks besides its sinks, at any history. + if window is not None: + blocks = min(target, int(layout.sink_blocks[lg]) + -(-int(window) // tpb)) + rows[int(group)] = rows.get(int(group), 0) + blocks + return rows + + +def slot_counts(layout: ManagerLayout, options: StagingOptions) -> Dict[int, int]: + """Slots per device pool group: ``max_fetches`` fetches, or ``max_bytes`` split by each group's + share of one fetch and rounded down to whole slots. ``ValueError`` if ``max_bytes`` is below + one fetch, so every group holds at least one fetch's rows.""" + rows = fetch_rows(layout, options.fetch_tokens) + weight = {g: rows[g] * int(layout.page_bytes[g]) for g in rows} + one = sum(weight.values()) + if one < 1: + raise ValueError("the KV cache manager has no pages to stage") + nbytes = options.max_fetches * one + if options.max_bytes is not None: + if options.max_bytes < one: + raise ValueError( + f"max_bytes of {options.max_bytes} is below one fetch of {options.fetch_tokens} " + f"tokens, which needs {one} bytes" + ) + nbytes = min(nbytes, options.max_bytes) + # Each group's share of the bytes is its share of one fetch, so a cap keeps every group's + # rows of one fetch; uncapped, a group gets exactly max_fetches times its rows. + return {g: (nbytes * weight[g] // one) // int(layout.page_bytes[g]) for g in rows} + + +@dataclass(frozen=True) +class Runs: + """The slots one lease holds: per pool group, ``count`` slots from ``start``.""" + + runs: Mapping[int, Tuple[int, int]] + + def slots(self, group: int) -> np.ndarray: + """``int64`` slot indices of ``group``'s run, empty when the lease holds none there.""" + start, count = self.runs.get(group, (0, 0)) + return np.arange(start, start + count, dtype=np.int64) + + +class _RunAllocator: + """First-fit allocator of contiguous slot runs with coalescing frees.""" + + def __init__(self, num_slots: int): + self.num_slots = num_slots + self._free: List[Tuple[int, int]] = [(0, num_slots)] if num_slots else [] + + def find(self, count: int) -> Optional[int]: + for start, length in self._free: + if length >= count: + return start + return None + + def take(self, start: int, count: int) -> None: + for i, (s, length) in enumerate(self._free): + if s <= start and start + count <= s + length: + pieces = [] + if start > s: + pieces.append((s, start - s)) + if start + count < s + length: + pieces.append((start + count, s + length - start - count)) + self._free[i : i + 1] = pieces + return + raise ValueError(f"slots [{start}, {start + count}) are not free") + + def overlaps_free(self, start: int, count: int) -> bool: + """Some slot of ``[start, start + count)`` is free already.""" + return any(s < start + count and start < s + length for s, length in self._free) + + def give(self, start: int, count: int) -> None: + if count == 0: + return + if start < 0 or count < 0 or start + count > self.num_slots: + raise ValueError( + f"slots [{start}, {start + count}) are outside the pool of {self.num_slots}" + ) + i = bisect.bisect_left(self._free, (start, 0)) + prev_end = self._free[i - 1][0] + self._free[i - 1][1] if i > 0 else -1 + if prev_end > start or (i < len(self._free) and self._free[i][0] < start + count): + raise ValueError(f"slots [{start}, {start + count}) freed twice") + self._free.insert(i, (start, count)) + merged: List[Tuple[int, int]] = [] + for s, length in self._free: + if merged and merged[-1][0] + merged[-1][1] == s: + merged[-1] = (merged[-1][0], merged[-1][1] + length) + else: + merged.append((s, length)) + self._free = merged + + @property + def free_slots(self) -> int: + return sum(length for _, length in self._free) + + +class Slots: + """First-come-first-served slot queue over every pool group. Not thread-safe: like the manager, + one thread at a time uses it.""" + + def __init__(self, num_slots: Mapping[int, int]) -> None: + self._num_slots: Dict[int, int] = {} + for g, count in num_slots.items(): + if int(count) < 0: + raise ValueError(f"pool group {g} cannot hold {count} slots") + self._num_slots[int(g)] = int(count) + self._alloc = {g: _RunAllocator(c) for g, c in self._num_slots.items()} + self._waiting: OrderedDict[int, Dict[int, int]] = OrderedDict() + self._next_ticket = 0 + + def num_slots(self, group: int) -> int: + """Slots ``group`` holds in all.""" + return self._num_slots[group] + + def check(self, counts: Mapping[int, int]) -> None: + """``ValueError`` if ``counts`` can never be granted: an unknown group, or more slots than a + group holds. Queues nothing.""" + self._wanted(counts) + + def ask(self, counts: Mapping[int, int]) -> int: + """Queue a request for ``counts`` slots per group and return its ticket; ``ValueError`` as + ``check``.""" + wanted = self._wanted(counts) + ticket = self._next_ticket + self._next_ticket += 1 + self._waiting[ticket] = wanted + return ticket + + def take(self, ticket: int) -> Optional[Runs]: + """Grant ``ticket`` if it is first in line and every group has a free run; else ``None``. + ``KeyError`` for a ticket that is not waiting.""" + if ticket not in self._waiting: + raise KeyError(f"staging ticket {ticket} is not waiting") + if next(iter(self._waiting)) != ticket: + return None + wanted = self._waiting[ticket] + starts = {} + for g, c in wanted.items(): + start = self._alloc[g].find(c) + if start is None: + return None + starts[g] = start + for g, c in wanted.items(): + self._alloc[g].take(starts[g], c) + del self._waiting[ticket] + return Runs({g: (starts[g], c) for g, c in wanted.items()}) + + def cancel(self, ticket: int) -> None: + """Withdraw a waiting ticket; an unknown or granted ticket is ignored.""" + self._waiting.pop(ticket, None) + + def give(self, runs: Runs) -> None: + """Return one lease's slots, all or none. ``ValueError`` for a run not held (freed twice or + never granted).""" + # Every run is checked before any is returned, so a bad one returns none. + for g, (start, count) in runs.runs.items(): + alloc = self._alloc.get(g) + if alloc is None: + raise ValueError(f"no staging slots for pool group {g}") + if count and (start < 0 or count < 0 or start + count > alloc.num_slots): + raise ValueError(f"slots [{start}, {start + count}) are outside pool group {g}") + if count and alloc.overlaps_free(start, count): + raise ValueError(f"slots [{start}, {start + count}) of pool group {g} freed twice") + for g, (start, count) in runs.runs.items(): + self._alloc[g].give(start, count) + + @property + def num_waiting(self) -> int: + """Tickets waiting in line.""" + return len(self._waiting) + + def free_slots(self, group: int) -> int: + """Slots of ``group`` not held by any lease.""" + return self._alloc[group].free_slots + + def _wanted(self, counts: Mapping[int, int]) -> Dict[int, int]: + """``counts`` without zero entries; ``ValueError`` if it can never be granted.""" + wanted = {} + for g, c in counts.items(): + g, c = int(g), int(c) + if c < 0: + raise ValueError(f"{c} staging slots asked of pool group {g}") + if c == 0: + continue + if g not in self._num_slots: + raise ValueError(f"no staging slots for pool group {g}") + if c > self._num_slots[g]: + raise ValueError( + f"{c} staging slots asked of pool group {g}, which has {self._num_slots[g]}" + ) + wanted[g] = c + return wanted diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_types.py b/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_types.py new file mode 100644 index 000000000000..ccadac766ff6 --- /dev/null +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache/sharing/_types.py @@ -0,0 +1,281 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""The public types and protocols of the lender. Numpy and the standard library only, so any +side can import them without a live KV cache manager.""" + +from __future__ import annotations + +import numbers +from dataclasses import dataclass +from typing import TYPE_CHECKING, NamedTuple, Optional, Protocol, Sequence, Tuple, runtime_checkable + +import numpy as np + +if TYPE_CHECKING: + from ...llm_request import LlmRequest + +# A row's name: a fixed-length opaque key. +_NAME_BYTES = 54 + + +def _freeze(array: np.ndarray) -> np.ndarray: + """``array`` marked read-only (a view when it is not already).""" + if array.flags.writeable: + array = array.view() + array.flags.writeable = False + return array + + +def _positive(name: str, value) -> int: + if isinstance(value, bool) or not isinstance(value, numbers.Integral): + raise TypeError(f"{name} must be an integer, got {type(value).__name__}") + if value < 1: + raise ValueError(f"{name} must be positive, got {value}") + return int(value) + + +@dataclass(frozen=True) +class StagingOptions: + """Staging in whole fetches: ``max_fetches`` of ``fetch_tokens`` tokens, capped by ``max_bytes`` + but never below one fetch, which always fits. A capacity budget, not concurrency: a lease takes + one contiguous run per pool group, and holes released leases leave can make it wait.""" + + fetch_tokens: int + max_fetches: int = 1 + max_bytes: Optional[int] = None + + def __post_init__(self): + object.__setattr__(self, "fetch_tokens", _positive("fetch_tokens", self.fetch_tokens)) + object.__setattr__(self, "max_fetches", _positive("max_fetches", self.max_fetches)) + if self.max_bytes is not None: + object.__setattr__(self, "max_bytes", _positive("max_bytes", self.max_bytes)) + + +@dataclass(frozen=True) +class Part: + """One device pool group's staging host region: ``nbytes`` bytes at ``address``, holding + ``slots`` slots of ``slot_bytes`` bytes, one row each. Instances laid out alike give it the + same ``name``.""" + + name: str + address: int + nbytes: int + slot_bytes: int + slots: int + + +@dataclass(frozen=True, eq=False) +class GroupRun: + """One layer group's read-only rows: row ``i`` is block ``ordinals[i]``. Staging rows also + carry an opaque 54-byte name (a store key: ``names[i].tobytes()``), their slot's address and + their part's index in ``StagingLender.parts``; in-place rows carry none of the three.""" + + layer_group: int + ordinals: np.ndarray + names: Optional[np.ndarray] = None + addresses: Optional[np.ndarray] = None + part: Optional[int] = None + + def __post_init__(self): + group = self.layer_group + if isinstance(group, bool) or not isinstance(group, numbers.Integral): + raise TypeError(f"layer_group must be an integer, got {type(group).__name__}") + if self.layer_group < 0: + raise ValueError(f"negative layer group {self.layer_group}") + object.__setattr__(self, "layer_group", int(self.layer_group)) + ordinals = _freeze(np.ascontiguousarray(self.ordinals, dtype=np.int64)) + if ordinals.ndim != 1: + raise ValueError(f"ordinals must be one-dimensional, got shape {ordinals.shape}") + object.__setattr__(self, "ordinals", ordinals) + placed = (self.names is not None, self.addresses is not None, self.part is not None) + if any(placed) and not all(placed): + raise ValueError("names, addresses and part are all given or all None") + if not all(placed): + return + names = np.asarray(self.names) + if names.dtype != np.uint8 or names.shape != (len(ordinals), _NAME_BYTES): + raise ValueError( + f"names must be uint8 ({len(ordinals)}, {_NAME_BYTES}), " + f"got {names.dtype} {names.shape}" + ) + addresses = np.asarray(self.addresses) + if not np.issubdtype(addresses.dtype, np.integer) or addresses.shape != ordinals.shape: + raise ValueError( + f"addresses must be integers of shape {ordinals.shape}, " + f"got {addresses.dtype} {addresses.shape}" + ) + if isinstance(self.part, bool) or not isinstance(self.part, numbers.Integral): + raise TypeError(f"part must be an integer, got {type(self.part).__name__}") + if self.part < 0: + raise ValueError(f"negative part {self.part}") + object.__setattr__(self, "names", _freeze(np.ascontiguousarray(names))) + object.__setattr__( + self, "addresses", _freeze(np.ascontiguousarray(addresses, dtype=np.int64)) + ) + object.__setattr__(self, "part", int(self.part)) + + def __len__(self) -> int: + return int(self.ordinals.shape[0]) + + def select(self, mask: np.ndarray) -> GroupRun: + """The rows where the boolean ``mask`` (one entry per row) is true, in order.""" + mask = np.asarray(mask) + if mask.dtype != np.bool_ or mask.shape != self.ordinals.shape: + raise ValueError( + f"mask must be bool of shape {self.ordinals.shape}, got {mask.dtype} {mask.shape}" + ) + return GroupRun( + self.layer_group, + self.ordinals[mask], + None if self.names is None else self.names[mask], + None if self.addresses is None else self.addresses[mask], + self.part, + ) + + +@dataclass(frozen=True, eq=False) +class RegionView: + """What a ready lease lends: at most one run per layer group. Backends may read its arrays on + their own threads until the lease is released.""" + + runs: Tuple[GroupRun, ...] + + def __post_init__(self): + runs = tuple(self.runs) + seen = set() + for run in runs: + if not isinstance(run, GroupRun): + raise TypeError(f"runs hold GroupRun, got {type(run).__name__}") + if run.layer_group in seen: + raise ValueError(f"layer group {run.layer_group} appears twice") + seen.add(run.layer_group) + object.__setattr__(self, "runs", runs) + + @property + def num_rows(self) -> int: + """Rows over all runs.""" + return sum(len(run) for run in self.runs) + + def row_masks(self, value: bool = False) -> Tuple[np.ndarray, ...]: + """One writable boolean array per run, filled with ``value``: the shape ``mark_arrived`` + takes.""" + return tuple(np.full(len(run), bool(value), dtype=bool) for run in self.runs) + + +class Readiness(NamedTuple): + """Resume at any ``p`` with ``restart_floor <= p <= usable_until``; ranks take the largest floor + and the smallest end, and an empty interval means drop the cache and compute from 0. Below the + floor a sliding window has released earlier blocks, so resuming there cannot continue.""" + + usable_until: int + restart_floor: int + + +@runtime_checkable +class Lease(Protocol): + """One lent range: poll until the view or ``failure``, mark a write once, release. Its methods + run only on the manager's thread; a backend's own threads read the view, access the memory it + points to, call no lease or lender method and tell the holder through their own channel.""" + + def poll(self) -> Optional[RegionView]: + """Does pending work. The view once ready (the same object each time); ``None`` while + pending and for good once failed. ``RuntimeError`` after release.""" + ... + + @property + def failure(self) -> Optional[str]: + """Why the lease will never be ready, once that is known; for logs.""" + ... + + def mark_arrived(self, masks: Sequence[np.ndarray]) -> None: + """Write leases, once, after ready, before or after release: one boolean mask per run, true + for rows that arrived whole. Required for staging writes, whose marked rows alone reach the + request's pages, in one batch: split long fetches. Optional in place; it checks shapes.""" + ... + + def release(self) -> None: + """The backend has stopped touching the lent memory. Legal in every state; later calls do + nothing.""" + ... + + +@runtime_checkable +class PartsHold(Protocol): + """A hold still open at the manager's shutdown keeps the staging memory until the process exits. + Take one before registering ``StagingLender.parts``; release it on the manager's thread once + deregistration is confirmed. The lender holds it, so dropping it unreleased keeps the memory.""" + + def release(self) -> None: + """The backend can no longer reach the parts. Runs only on the manager's thread, like every + lender and lease method; legal in every state, later calls do nothing.""" + ... + + +# TODO: a naming entry that allocates nothing, for a flow that looks up remote hits before it +# reserves pages for them; a pure addition beside lend_read and lend_write. +@runtime_checkable +class StagingLender(Protocol): + """Relays whole blocks through host staging slots, first come, first served. Lender and lease + methods run only on the manager's thread; a backend's own threads read a view, access the memory + it points to, call no lease or lender method, and tell the holder through their own channel.""" + + @property + def parts(self) -> Tuple[Part, ...]: + """The staging host regions, one per device pool group, fixed for the lender's life and + registrable once. Freed at the manager's shutdown unless an unreleased lease, an unreleased + hold or a slot lost to a failed copy keeps them until the process exits.""" + ... + + def hold_parts(self) -> PartsHold: + """A new hold on the staging memory, taken before registering ``parts``. After the + manager's shutdown the hold is inert: the memory was freed or kept then.""" + ... + + def lend_read(self, request: LlmRequest, start: int, end: int) -> Lease: + """A copy of the committed blocks ``[start, end)``; ``ValueError`` for bad bounds, an end + past the committed tokens or a misfit range. Its outcome is this rank's own, failed at the + call or later on this rank's state; the caller combines every rank's outcome.""" + ... + + def lend_write(self, request: LlmRequest, start: int, end: int) -> Lease: + """Empty slots for blocks ``[start, end)``, the cache grown to ``end``. ``ValueError``: bad + bounds, a start in the committed blocks, an end below a window's history, a misfit range, + an unsettled fetch. Fails as ``lend_read``, also on no free pages or SWA scratch reuse.""" + ... + + def readiness(self, request: LlmRequest) -> Optional[Readiness]: + """``None`` while a fetch into the request is unsettled, else where it may resume. Copies + queue on the manager's stream, serial with the forward on the GPU; with page-locked staging + no lender call waits for them on the CPU. ``ValueError`` if the request has no KV cache.""" + ... + + +@runtime_checkable +class InPlaceLender(Protocol): + """Lends a request's own device pages, addressed by the caller's own page-table code. Lender and + lease methods run only on the manager's thread; a backend's own threads access the lent memory, + call no lease or lender method, and tell the holder through their own channel.""" + + def lend_read(self, request: LlmRequest, start: int, end: int) -> Lease: + """The paged blocks ``[start, end)`` touches, ready at once; ``ValueError`` for a bad range. + Its outcome is this rank's own; the caller combines every rank's outcome, and ensures work + the manager's stream queued that still writes those pages completed.""" + ... + + def lend_write(self, request: LlmRequest, start: int, end: int) -> Lease: + """As ``lend_read`` for writing; a block without a page fails it too, and all work the + manager's stream queued for those pages has completed. While either is lent the caller keeps + the request unscheduled, unsuspended, unshrunk and its window still, and owns validity.""" + ... diff --git a/tests/integration/test_lists/test-db/l0_a10.yml b/tests/integration/test_lists/test-db/l0_a10.yml index 53ac082dcb0f..0eef8774576d 100644 --- a/tests/integration/test_lists/test-db/l0_a10.yml +++ b/tests/integration/test_lists/test-db/l0_a10.yml @@ -39,6 +39,7 @@ l0_a10: - unittest/_torch/disaggregation/test_disagg_index_mapper_early_release.py - unittest/_torch/executor/kv_cache/test_kv_cache_compression_manager.py - unittest/_torch/executor/kv_cache/test_kv_cache_v2_capacity_only.py + - unittest/_torch/executor/kv_cache/sharing - unittest/_torch/executor/test_kv_cache_layout.py - unittest/_torch/executor/test_kv_connector_v2_prefix_real_manager.py - unittest/_torch/executor/test_mooncake_store_cli.py diff --git a/tests/unittest/_torch/executor/kv_cache/sharing/conftest.py b/tests/unittest/_torch/executor/kv_cache/sharing/conftest.py new file mode 100644 index 000000000000..410e3b7b7fd4 --- /dev/null +++ b/tests/unittest/_torch/executor/kv_cache/sharing/conftest.py @@ -0,0 +1,954 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Real KV cache managers, requests and byte oracles for the lender tests. The oracles bypass the +code under test: device pages through ``impl.pool_group_descs``, staging bytes through raw host +addresses, block keys through ``sequence_to_blockchain_keys``. Nothing here touches CUDA at import.""" + +from __future__ import annotations + +import ctypes +import gc +import hashlib +import inspect +import threading +import time +from contextlib import contextmanager +from types import SimpleNamespace +from typing import Callable, Dict, Iterator, List, Optional, Sequence, Tuple +from unittest.mock import patch + +import numpy as np +import pytest +import torch + +TPB = 32 +WINDOW = 64 +MAX_SEQ_LEN = 256 +SENTINEL = 0xEE +# The smallest pool the harness manager accepts: few enough pages that freed ones get reused. +POOL_TOKENS = 256 + +_SHARING = "tensorrt_llm._torch.pyexecutor.kv_cache.sharing" + + +class _RanksAgree: + """The collectives of a TP manager built alone in this process: every rank agrees with it.""" + + local_world_size = 1 + + @staticmethod + def allreduce(value, op=None): + return value + + +@contextmanager +def _collectives_for(mapping): + """Stand-in collectives while a multi-rank manager is built in a single process.""" + if mapping is None or mapping.world_size == 1: + yield + return + from tensorrt_llm._torch.distributed.communicator import Distributed + + with patch.object(Distributed, "get", return_value=_RanksAgree()): + yield + + +def make_manager( + *, + windows: Optional[List[int]] = None, + max_tokens: int = 2048, + tokens_per_block: int = TPB, + num_layers: int = 2, + num_kv_heads: int = 4, + head_dim: int = 64, + mapping=None, + kv_cache_type=None, + dtype=None, + enable_block_reuse: bool = True, + swa_scratch_reuse: bool = False, + max_batch_size: int = 4, + **kv_cache_config, +): + """A small real ``KVCacheManagerV2`` on its own execution stream, FP16 unless ``dtype``; further + keywords go to its ``KvCacheConfig``.""" + import tensorrt_llm.bindings + from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import KVCacheManagerV2 + from tensorrt_llm.llmapi.llm_args import KvCacheConfig + from tensorrt_llm.mapping import Mapping + + batch_manager = tensorrt_llm.bindings.internal.batch_manager + config = dict(max_tokens=max_tokens, enable_block_reuse=enable_block_reuse) + if swa_scratch_reuse: + config["enable_swa_scratch_reuse"] = True + if windows is not None: + config["max_attention_window"] = windows + config.update(kv_cache_config) + mapping = mapping or Mapping(world_size=1, tp_size=1, rank=0) + with _collectives_for(mapping): + return KVCacheManagerV2( + kv_cache_config=KvCacheConfig(**config), + kv_cache_type=kv_cache_type or batch_manager.CacheType.SELF, + num_layers=num_layers, + num_kv_heads=num_kv_heads, + head_dim=head_dim, + tokens_per_block=tokens_per_block, + max_seq_len=MAX_SEQ_LEN, + max_batch_size=max_batch_size, + mapping=mapping, + dtype=dtype or tensorrt_llm.bindings.DataType.HALF, + vocab_size=32000, + execution_stream=torch.cuda.Stream(), + ) + + +def _make_deepseek_v4_manager(mapping=None): + """A four-layer DeepSeek-V4 manager: one layer of each compression kind, FP8, one KV head.""" + from tensorrt_llm._torch.attention.backends.sparse.deepseek_v4 import DeepseekV4CacheManager + from tensorrt_llm.bindings import DataType + from tensorrt_llm.bindings.internal.batch_manager import CacheType + from tensorrt_llm.llmapi.llm_args import DeepSeekV4SparseAttentionConfig, KvCacheConfig + from tensorrt_llm.mapping import Mapping + + compress_ratios = [1, 4, 128, 4] + mapping = mapping or Mapping(world_size=1, rank=0, tp_size=1, pp_size=1) + with _collectives_for(mapping): + return DeepseekV4CacheManager( + kv_cache_config=KvCacheConfig( + enable_block_reuse=True, max_tokens=4096, event_buffer_max_size=0 + ), + kv_cache_type=CacheType.SELFKONLY, + num_layers=len(compress_ratios), + num_kv_heads=1, + head_dim=512, + tokens_per_block=128, + max_seq_len=1024, + max_batch_size=2, + max_input_len=1024, + mapping=mapping, + dtype=DataType.FP8, + compressor_dtype=DataType.FLOAT, + vocab_size=129280, + max_num_tokens=2 * 1025, + sparse_attn_config=DeepSeekV4SparseAttentionConfig( + index_head_dim=128, window_size=128, compress_ratios=compress_ratios + ), + execution_stream=torch.cuda.Stream(), + ) + + +def _make_hybrid_manager(): + """Nemotron-H shaped, few layers: Mamba2 state at layers 0, 2 and 4, attention at layer 3.""" + from tensorrt_llm._torch.pyexecutor.kv_cache.mamba_cache_manager import ( + MambaHybridCacheManagerV2, + ) + from tensorrt_llm.bindings import DataType + from tensorrt_llm.bindings.internal.batch_manager import CacheType + from tensorrt_llm.llmapi.llm_args import KvCacheConfig, MambaStateConfig + from tensorrt_llm.mapping import Mapping + + pattern = "M-M*M-" + mamba_mask = [c == "M" for c in pattern] + attn_mask = [c == "*" for c in pattern] + with _collectives_for(Mapping(world_size=1, rank=0, tp_size=1)): + return MambaHybridCacheManagerV2( + mamba_d_state=128, + mamba_d_conv=4, + mamba_num_heads=128, + mamba_n_groups=8, + mamba_head_dim=80, + mamba_num_layers=sum(mamba_mask), + mamba_layer_mask=mamba_mask, + mamba_cache_dtype=torch.bfloat16, + mamba_ssm_cache_dtype=torch.float32, + kv_cache_config=KvCacheConfig( + max_tokens=2048, + enable_block_reuse=True, + mamba_state_config=MambaStateConfig(periodic_snapshot_interval=256), + ), + kv_cache_type=CacheType.SELF, + num_layers=sum(attn_mask), + num_kv_heads=8, + head_dim=128, + tokens_per_block=32, + max_seq_len=1024, + max_batch_size=2, + mapping=Mapping(world_size=1, rank=0, tp_size=1), + layer_mask=attn_mask, + vocab_size=1024, + dtype=DataType.BF16, + ) + + +@contextmanager +def _managed(factory: Callable[[], object]) -> Iterator[object]: + """A manager that is shut down, with every request's cache freed, when the block exits.""" + torch.cuda.init() + gc.collect() + torch.cuda.empty_cache() + mgr = factory() + try: + yield mgr + finally: + stream = getattr(mgr, "_stream", None) + if stream is not None: + stream.synchronize() + mgr.shutdown() + del mgr + gc.collect() + torch.cuda.empty_cache() + + +@pytest.fixture +def real_manager(): + """``with real_manager(windows=None, max_tokens=2048, ...) as mgr:`` a small ``KVCacheManagerV2`` + with block reuse on and its own execution stream, shut down when the block exits.""" + + def factory(**kwargs): + return _managed(lambda: make_manager(**kwargs)) + + return factory + + +@contextmanager +def _host_tiered(factory: Callable[[], object]) -> Iterator[object]: + """``_managed``, checked to have the automatic host tier below its GPU pool.""" + with _managed(factory) as mgr: + tiers = [str(tier) for tier in mgr.impl.cache_tier_list] + assert len(tiers) == 2 and "HOST" in tiers[1], f"no host tier below the GPU pool: {tiers}" + yield mgr + + +@pytest.fixture +def host_tier_manager(): + """``with host_tier_manager(**kwargs) as mgr:`` like ``real_manager`` on a pool of + ``POOL_TOKENS``, with its automatic host tier checked present. Index slots for many one-block + requests and resume allowed up to a full pool let other requests push pages to host and back.""" + + def factory(**kwargs): + config = dict(max_tokens=POOL_TOKENS, max_batch_size=64, max_util_for_resume=1.0) + config.update(kwargs) + return _host_tiered(lambda: make_manager(**config)) + + return factory + + +@pytest.fixture +def deepseek_v4_manager(): + """``with deepseek_v4_manager(mapping=None) as mgr:`` a small ``DeepseekV4CacheManager``; + callers skip before Blackwell.""" + + def factory(mapping=None): + return _managed(lambda: _make_deepseek_v4_manager(mapping)) + + return factory + + +@pytest.fixture +def hybrid_manager(): + """``with hybrid_manager() as mgr:`` a manager whose layer groups include recurrent state.""" + + def factory(): + return _managed(_make_hybrid_manager) + + return factory + + +# -- requests --------------------------------------------------------------------------------- + + +def make_request(request_id: int, tokens: Sequence[int], **kwargs): + from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, SamplingConfig + + return LlmRequest( + request_id=request_id, + max_new_tokens=4, + input_tokens=list(tokens), + sampling_config=SamplingConfig(1), + is_streaming=False, + **kwargs, + ) + + +def kv(mgr, request): + """The request's runtime cache, or ``None``: read straight from the manager's map.""" + return mgr.kv_cache_map.get(request.py_request_id) + + +def closed(kv_cache) -> bool: + return kv_cache.status == kv_cache.Status.CLOSED + + +def chain_keys(kv_cache, tokens: Sequence) -> List[bytes]: + """The radix keys of every whole block of ``tokens``, computed independently of the lender.""" + from tensorrt_llm.runtime.kv_cache_manager_v2 import sequence_to_blockchain_keys + + tpb = kv_cache.tokens_per_block + keys = [] + for i, (block, key) in enumerate( + sequence_to_blockchain_keys(tpb, kv_cache.reuse_scope, list(tokens)) + ): + if i and len(block) == tpb: + keys.append(bytes(key)) + return keys + + +def fake_page(lg: int, key: bytes, nbytes: int) -> bytes: + """Deterministic page content named by (layer group, block key): equal tokens, equal bytes.""" + seed = int.from_bytes(hashlib.sha256(bytes([lg]) + key).digest()[:8], "little") + return np.random.default_rng(seed).integers(0, 256, nbytes, dtype=np.uint8).tobytes() + + +def pages(kv_cache, lg: int) -> List[int]: + return [int(p) for p in kv_cache.get_aggregated_page_indices(lg, valid_only=False)] + + +def num_layer_groups(mgr) -> int: + return len(mgr.impl.layer_grouping) + + +def pool_group_of(mgr) -> List[int]: + """Device pool group of each layer group.""" + return [int(g) for g in mgr.impl.get_life_cycle_pool_group_indices()] + + +def pool_group_ids(mgr) -> List[int]: + """Device pool groups in index order: the order of a staging lender's parts.""" + return sorted(int(pg.pool_group_index) for pg in mgr.impl.pool_group_descs) + + +def windows(mgr) -> Tuple[Optional[int], ...]: + """Each layer group's sliding window in tokens, ``None`` for full attention.""" + out = [] + for lc in mgr._life_cycle_by_layer_group(): + window = getattr(lc, "window_size", None) + out.append(None if window is None or window >= MAX_SEQ_LEN else int(window)) + return tuple(out) + + +def stale_blocks(mgr, lg: int, history: int) -> Tuple[int, int]: + beg, end = mgr._stale_block_range(lg, history) + return int(beg), int(end) + + +class DevicePages: + """Device pages addressed by (layer group, slot), straight from ``impl.pool_group_descs``; all + work runs on the manager's stream.""" + + def __init__(self, mgr): + from tensorrt_llm._utils import TensorWrapper, convert_to_torch_tensor + + self._mgr = mgr + self._group_of = pool_group_of(mgr) + self._pools: Dict[int, List[torch.Tensor]] = {} + for pg in mgr.impl.pool_group_descs: + ordered = sorted(pg.pools, key=lambda p: int(p.pool_index)) + self._pools[int(pg.pool_group_index)] = [ + convert_to_torch_tensor( + TensorWrapper( + int(p.base_address), + torch.uint8, + shape=(int(pg.num_slots), int(p.slot_bytes)), + ) + ) + for p in ordered + ] + # Loading the fill kernel waits for the whole device, so load it before any stream is held. + torch.empty(1, dtype=torch.uint8, device="cuda").fill_(0) + torch.cuda.synchronize() + + def group_page_bytes(self, group: int) -> int: + return sum(t.shape[1] for t in self._pools[group]) + + def page_bytes(self, lg: int) -> int: + return self.group_page_bytes(self._group_of[lg]) + + def read(self, lg: int, slot: int) -> bytes: + self._mgr._stream.synchronize() + tensors = self._pools[self._group_of[lg]] + return b"".join(t[slot].cpu().numpy().tobytes() for t in tensors) + + def write(self, lg: int, slot: int, data: bytes) -> None: + stream = self._mgr._stream + src = torch.frombuffer(bytearray(data), dtype=torch.uint8) + offset = 0 + with torch.cuda.stream(stream): + for t in self._pools[self._group_of[lg]]: + width = t.shape[1] + t[slot].copy_(src[offset : offset + width].to(t.device)) + offset += width + stream.synchronize() + + def fill_async(self, lg: int, slot: int, byte: int) -> None: + """Fill a page on the manager's stream without waiting: ordered after work already queued.""" + with torch.cuda.stream(self._mgr._stream): + for t in self._pools[self._group_of[lg]]: + t[slot].fill_(byte) + + +def page(mgr, request, lg: int, ordinal: int) -> bytes: + return DevicePages(mgr).read(lg, pages(kv(mgr, request), lg)[ordinal]) + + +def write_fake_kv(mgr, request, first_block: int = 0) -> None: + """Fill every page the request holds from ``first_block`` on, as a forward pass would.""" + kv_cache = kv(mgr, request) + dev = DevicePages(mgr) + keys = chain_keys(kv_cache, request.get_tokens(0)) + for lg in range(num_layer_groups(mgr)): + for ordinal, slot in enumerate(pages(kv_cache, lg)): + if ordinal < first_block or slot < 0 or ordinal >= len(keys): + continue + dev.write(lg, slot, fake_page(lg, keys[ordinal], dev.page_bytes(lg))) + + +def fill_sentinel(mgr, request, first_block: int = 0) -> None: + """Mark the request's pages from ``first_block`` on so a missing copy shows.""" + kv_cache = kv(mgr, request) + dev = DevicePages(mgr) + for lg in range(num_layer_groups(mgr)): + for ordinal, slot in enumerate(pages(kv_cache, lg)): + if ordinal >= first_block and slot >= 0: + dev.write(lg, slot, bytes([SENTINEL]) * dev.page_bytes(lg)) + + +def prefill(mgr, request) -> None: + """Run a whole prompt the way the executor does: allocate, compute (fake KV), advance, then + ``update_context_resources``, which commits the blocks.""" + from tensorrt_llm._torch.pyexecutor.scheduler import ScheduledRequests + + assert mgr.prepare_context(request) + kv_cache = kv(mgr, request) + # Real pages for every block written below; scratch slots would be overwritten. + kv_cache.enable_swa_scratch_reuse = False + first = request.context_current_position // mgr.tokens_per_block + assert mgr.resize_context(request, request.context_remaining_length) + write_fake_kv(mgr, request, first) + request.move_to_next_context_chunk() + batch = ScheduledRequests() + batch.append_context_request(request) + mgr.update_context_resources(batch) + assert request.context_remaining_length == 0 + assert kv_cache.num_committed_tokens == request.prompt_len + + +def published(mgr, request_id: int, tokens: Sequence[int], **kwargs): + """A request that computed ``tokens``: its blocks are committed.""" + request = make_request(request_id, tokens, **kwargs) + prefill(mgr, request) + return request + + +def admitted(mgr, request_id: int, tokens: Sequence[int]): + """A request admitted to fetch its prompt: a cache holding its local match and nothing more, + scratch reuse off as a fetch target needs. The lender grows it.""" + request = make_request(request_id, tokens) + assert mgr.prepare_context(request) + kv(mgr, request).enable_swa_scratch_reuse = False + return request + + +def host_bytes(address: int, length: int) -> bytes: + return ctypes.string_at(address, length) + + +def digest(value): + """``value`` with every byte string replaced by its length and hash, or its one repeated byte: + a failed comparison then names the entries that differ without diffing page contents.""" + if isinstance(value, (bytes, bytearray)): + if value and value.count(value[:1]) == len(value): + return f"{len(value)} x {value[0]:#04x}" + return f"{len(value)} B sha256 {hashlib.sha256(value).hexdigest()[:16]}" + if isinstance(value, dict): + return {key: digest(item) for key, item in value.items()} + if isinstance(value, (list, tuple)): + return [digest(item) for item in value] + return value + + +def stage(lender, view, byte: int) -> None: + """Write ``byte`` over every slot of a staging view, as a backend's transfer would.""" + for run in view.runs: + length = lender.parts[run.part].slot_bytes + for address in run.addresses.tolist(): + ctypes.memset(address, byte, length) + + +def relay(src_lender, src_view, dst_lender, dst_view, deliver=None) -> Tuple[np.ndarray, ...]: + """A backend between two staging lenders: each destination row takes the slot of the source row + of the same name. ``deliver(run_index, row)`` picks the rows (default all); returns the masks.""" + by_name = {} + for run in src_view.runs: + length = src_lender.parts[run.part].slot_bytes + for name, address in zip(run.names, run.addresses.tolist()): + by_name[name.tobytes()] = (address, length) + masks = [] + for i, run in enumerate(dst_view.runs): + mask = np.zeros(len(run), dtype=bool) + for row, (name, address) in enumerate(zip(run.names, run.addresses.tolist())): + if deliver is not None and not deliver(i, row): + continue + src, length = by_name[name.tobytes()] + assert length == dst_lender.parts[run.part].slot_bytes + ctypes.memmove(address, src, length) + mask[row] = True + masks.append(mask) + return tuple(masks) + + +# -- streams ---------------------------------------------------------------------------------- + + +class _Gate: + def __init__(self, flag: torch.Tensor): + self._flag = flag + self.opened = False + self.tripped_by_watchdog = False + + def open(self) -> None: + self._flag.numpy()[0] = 1 + self.opened = True + + +def _wait_on_flag(stream: torch.cuda.Stream) -> _Gate: + from cuda.bindings import driver + + flag = torch.zeros(1, dtype=torch.int32, pin_memory=True) + (err,) = driver.cuStreamWaitValue32( + driver.CUstream(stream.cuda_stream), + driver.CUdeviceptr(flag.data_ptr()), + 1, + driver.CUstreamWaitValue_flags.CU_STREAM_WAIT_VALUE_GEQ, + ) + if err == driver.CUresult.CUDA_ERROR_NOT_SUPPORTED: + pytest.skip("this device cannot hold a stream on a host flag") + assert err == driver.CUresult.CUDA_SUCCESS, err + return _Gate(flag) + + +@contextmanager +def held_stream(stream: torch.cuda.Stream, strict: bool = True, watchdog_seconds: float = 120.0): + """Hold ``stream`` on a host flag until ``gate.open()``: work queued meanwhile does not start. + ``strict`` makes a torch host sync raise; a watchdog opens the gate if the test never does, and + the block then fails, so a host wait on the stream cannot pass unnoticed.""" + gate = _wait_on_flag(stream) + + def trip(): + gate.tripped_by_watchdog = True + gate.open() + + watchdog = threading.Timer(watchdog_seconds, trip) + watchdog.start() + previous = torch.cuda.get_sync_debug_mode() + if strict: + torch.cuda.set_sync_debug_mode("error") + try: + yield gate + finally: + torch.cuda.set_sync_debug_mode(previous) + gate.open() + watchdog.cancel() + watchdog.join() + stream.synchronize() + assert not gate.tripped_by_watchdog, "the stream stayed held: something waited on it" + + +@contextmanager +def gated_stream(stream: torch.cuda.Stream, open_after: float): + """Hold ``stream`` until a timer opens it after ``open_after`` seconds; a host wait on the + stream then simply lasts that long. ``gate.opened`` tells whether it has.""" + gate = _wait_on_flag(stream) + timer = threading.Timer(open_after, gate.open) + timer.start() + try: + yield gate + finally: + timer.cancel() + gate.open() + timer.join() + stream.synchronize() + + +@contextmanager +def events_report_done(): + """Every CUDA event reports its work done while the block runs, finished or not.""" + original = torch.cuda.Event.query + torch.cuda.Event.query = lambda self: True + try: + yield + finally: + torch.cuda.Event.query = original + + +# -- pool pressure ---------------------------------------------------------------------------- + + +def pool_pages(mgr) -> int: + (group,) = mgr.impl.pool_group_descs + return int(group.num_slots) + + +def warm_tier_moves(mgr) -> None: + """Move one committed block to host and back. The first such move in a process loads the kernels + that copy between tiers, which waits for the whole device; a held stream would stall it.""" + tokens = list(range(40_000, 40_000 + TPB + 1)) + mgr.free_resources(published(mgr, 900, tokens)) + others = Requests(mgr) + try: + while others.allocate(1, chunk=1): # until the pool is full: the block went to host + pass + finally: + others.free() + back = make_request(901, tokens) + assert mgr.prepare_context(back) + assert kv(mgr, back).num_committed_tokens == TPB, "the block did not come back from host" + mgr.free_resources(back) + torch.cuda.synchronize() + + +def tier_used(mgr, level: int) -> int: + """Slots in use at cache tier ``level`` (0 the GPU, 1 the host) over every pool group, from + the manager's storage statistics.""" + return sum(int(s.total) - int(s.free) for s in mgr.impl.get_storage_statistics(level)) + + +class Requests: + """Other requests' allocations: private pages taken straight from the pool, sentinel-filled, + a few blocks per request (a request holds at most ``MAX_SEQ_LEN`` tokens).""" + + def __init__(self, mgr, dev: Optional[DevicePages] = None): + # ``dev`` built ahead lets ``fill=False`` run while the stream is held: building the page + # views waits on the stream. + self._mgr = mgr + self._dev = dev + self._next = 100 + self.held = [] + + def allocate(self, blocks: int, fill: bool = True, chunk: Optional[int] = None) -> bool: + """Take ``blocks`` pages, ``chunk`` blocks per request. ``False`` if the pool could not give + them all; what it gave stays held. ``fill=False`` fills on the stream without waiting.""" + mgr = self._mgr + most = chunk or MAX_SEQ_LEN // TPB - 1 + while blocks: + chunk = min(blocks, most) + rid = self._next + self._next += 1 + first = 50_000 + 1_000 * rid # never matches another request's blocks + request = make_request(rid, list(range(first, first + chunk * TPB))) + if not ( + mgr.prepare_context(request) + and mgr.resize_context(request, request.context_remaining_length) + ): + mgr.free_resources(request) + return False + if fill: + fill_sentinel(mgr, request) + else: + dev = self._dev or DevicePages(mgr) + for lg in range(num_layer_groups(mgr)): + for slot in pages(kv(mgr, request), lg): + if slot >= 0: + dev.fill_async(lg, slot, SENTINEL) + self.held.append(request) + blocks -= chunk + return True + + def pages(self) -> set: + return {p for r in self.held for p in pages(kv(self._mgr, r), 0) if p >= 0} + + def free_all_but(self, keep: set) -> None: + """Free the requests holding none of the pages in ``keep``.""" + kept = [] + for request in self.held: + if keep & {p for p in pages(kv(self._mgr, request), 0) if p >= 0}: + kept.append(request) + else: + self._mgr.free_resources(request) + self.held = kept + + def free(self) -> None: + for request in self.held: + self._mgr.free_resources(request) + self.held = [] + + +def overwrite_free_pages(mgr) -> int: + """Other requests take every free page one block at a time, fill it with ``SENTINEL`` and free + it again, so whatever a freed page held is gone. Returns how many pages they took.""" + others = Requests(mgr) + try: + while others.allocate(1, chunk=1): + pass + return len(others.held) + finally: + others.free() + + +def taken_by_others(mgr, lent: set) -> set: + """Other requests ask for every page the pool has free besides ``lent``, then for one more. + Returns the pages of ``lent`` they got; they are freed again before returning.""" + others = Requests(mgr) + try: + assert others.allocate(pool_pages(mgr) - len(lent)) + others.allocate(1) + return others.pages() & lent + finally: + others.free() + + +def whole_pool_goes_to_others(mgr, lent: set) -> bool: + """Whether other requests can take every page of the pool, ``lent`` included.""" + others = Requests(mgr) + try: + return others.allocate(pool_pages(mgr)) and lent <= others.pages() + finally: + others.free() + + +class ShutdownSpy: + """Stands in for ``mgr.impl``, recording whether the manager shut it down.""" + + def __init__(self, impl): + self._impl = impl + self.shut_down = False + + def shutdown(self): + self.shut_down = True + self._impl.shutdown() + + def __getattr__(self, name): + return getattr(self._impl, name) + + +class StatsSpy: + """Stands in for ``mgr.impl``, recording each clear of a stats exclusion and whether the cache + was closed by then.""" + + def __init__(self, impl, kv_cache): + self._impl = impl + self._kv = kv_cache + self.cleared = [] + + def clear_stats_excluded(self, request_id): + self.cleared.append((request_id, closed(self._kv))) + self._impl.clear_stats_excluded(request_id) + + def __getattr__(self, name): + return getattr(self._impl, name) + + +# -- memory kept until exit ------------------------------------------------------------------- + + +def retained() -> Tuple[object, ...]: + """What the lender keeps until the process exits, through its test hook.""" + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + return _lender._retained() + + +def staging_memory(parts) -> List[object]: + """The kept host allocations holding any of ``parts``, compared by address.""" + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + found = [] + for obj in retained(): + if not isinstance(obj, _lender._HostMemory): + continue + begin, end = obj.address, obj.address + obj.nbytes + if any(begin <= part.address < end for part in parts): + found.append(obj) + return found + + +# -- the manager's host page-index buffer ------------------------------------------------------ + +CANARY = 0x5A5A5A5A + + +class IndexRow(SimpleNamespace): + """Where a request's cache writes its base page indices of one pool in the manager's host + page-index buffer: the buffer's address, size and pinning, and the row's int32 offset in it.""" + + +def index_row(mgr, request, pool: int = 0, beam: int = 0) -> IndexRow: + buf = mgr.host_kv_cache_block_offsets + index = mgr.index_mapper.get_index(int(request.py_request_id)) + row = buf[pool, index * mgr.max_beam_width + beam, 0] + return IndexRow( + address=buf.data_ptr(), + nbytes=buf.numel() * buf.element_size(), + pinned=buf.is_pinned(), + offset=(row.data_ptr() - buf.data_ptr()) // 4, + length=row.numel(), + values=row.tolist(), + ) + + +def reclaim(row: IndexRow, tries: int = 64) -> Optional[torch.Tensor]: + """A new host int32 tensor at a freed buffer's address, every element ``CANARY``; ``None`` if + the allocator hands that address to none of ``tries`` same-size allocations.""" + held = [] + for _ in range(tries): + t = torch.empty(row.nbytes // 4, dtype=torch.int32, pin_memory=row.pinned) + if t.data_ptr() == row.address: + t.fill_(CANARY) + return t + held.append(t) + return None + + +def canary_written(canary: Optional[torch.Tensor], row: IndexRow) -> List[Tuple[int, int]]: + """``(cell, value)`` of each cell of ``row`` in ``canary`` that is no longer ``CANARY``.""" + if canary is None: + return [] + cells = canary[row.offset : row.offset + row.length].numpy() + return [(int(i), int(cells[i])) for i in np.nonzero(cells != np.int32(CANARY))[0]] + + +def staging_kept(parts) -> bool: + """Whether every part lies in memory the lender keeps.""" + kept = staging_memory(parts) + return all( + any(m.address <= p.address and p.address + p.nbytes <= m.address + m.nbytes for m in kept) + for p in parts + ) + + +def pinned_range(address: int) -> Optional[Tuple[int, int]]: + """``(start, nbytes)`` of the page-locked host allocation the driver records at ``address``; + ``None`` when there is none, as after it was freed.""" + from cuda.bindings import driver + + attribute = driver.CUpointer_attribute + err, kind = driver.cuPointerGetAttribute(attribute.CU_POINTER_ATTRIBUTE_MEMORY_TYPE, address) + if err != driver.CUresult.CUDA_SUCCESS: + return None + assert kind == driver.CUmemorytype.CU_MEMORYTYPE_HOST, kind + err, start = driver.cuPointerGetAttribute( + attribute.CU_POINTER_ATTRIBUTE_RANGE_START_ADDR, address + ) + assert err == driver.CUresult.CUDA_SUCCESS, err + err, nbytes = driver.cuPointerGetAttribute(attribute.CU_POINTER_ATTRIBUTE_RANGE_SIZE, address) + assert err == driver.CUresult.CUDA_SUCCESS, err + return int(start), int(nbytes) + + +def lender_warnings(monkeypatch) -> List[str]: + """Collects the warnings the lender module logs from now on.""" + from tensorrt_llm.logger import logger + + seen: List[str] = [] + + def spy(original): + def warn(*args, **kwargs): + caller = inspect.currentframe().f_back + if caller is not None and caller.f_globals.get("__name__", "").startswith(_SHARING): + seen.append(" ".join(str(a) for a in args)) + return original(*args, **kwargs) + + return warn + + for name in ("warning", "warning_once"): + monkeypatch.setattr(logger, name, spy(getattr(logger, name))) + return seen + + +# -- threads ---------------------------------------------------------------------------------- + + +def on_thread(call: Callable[[], object], name: str = "not-the-caller") -> dict: + """Run ``call`` on a fresh thread, joined before returning; what it returned or raised.""" + box: dict = {} + + def run(): + try: + box["value"] = call() + except Exception as error: # handed back to the test + box["error"] = error + + thread = threading.Thread(target=run, name=name) + thread.start() + thread.join(120) + assert not thread.is_alive() + box["thread"] = thread.ident + return box + + +def wait_until(predicate: Callable[[], bool], seconds: float = 30.0) -> bool: + deadline = time.monotonic() + seconds + while not predicate(): + if time.monotonic() > deadline: + return False + time.sleep(0.01) + return True + + +_KIT = SimpleNamespace( + TPB=TPB, + WINDOW=WINDOW, + MAX_SEQ_LEN=MAX_SEQ_LEN, + SENTINEL=SENTINEL, + POOL_TOKENS=POOL_TOKENS, + CANARY=CANARY, + DevicePages=DevicePages, + Requests=Requests, + ShutdownSpy=ShutdownSpy, + StatsSpy=StatsSpy, + admitted=admitted, + canary_written=canary_written, + chain_keys=chain_keys, + closed=closed, + digest=digest, + events_report_done=events_report_done, + fill_sentinel=fill_sentinel, + gated_stream=gated_stream, + held_stream=held_stream, + host_bytes=host_bytes, + index_row=index_row, + kv=kv, + lender_warnings=lender_warnings, + make_manager=make_manager, + make_request=make_request, + num_layer_groups=num_layer_groups, + on_thread=on_thread, + overwrite_free_pages=overwrite_free_pages, + page=page, + pages=pages, + pinned_range=pinned_range, + pool_group_ids=pool_group_ids, + pool_group_of=pool_group_of, + pool_pages=pool_pages, + prefill=prefill, + published=published, + reclaim=reclaim, + relay=relay, + retained=retained, + stage=stage, + stale_blocks=stale_blocks, + staging_kept=staging_kept, + staging_memory=staging_memory, + taken_by_others=taken_by_others, + tier_used=tier_used, + warm_tier_moves=warm_tier_moves, + wait_until=wait_until, + whole_pool_goes_to_others=whole_pool_goes_to_others, + windows=windows, +) + + +@pytest.fixture +def kit(): + """The helpers and oracles above, for test modules that cannot import this file by name.""" + return _KIT diff --git a/tests/unittest/_torch/executor/kv_cache/sharing/test_host_tier.py b/tests/unittest/_torch/executor/kv_cache/sharing/test_host_tier.py new file mode 100644 index 000000000000..dcbff39f0d39 --- /dev/null +++ b/tests/unittest/_torch/executor/kv_cache/sharing/test_host_tier.py @@ -0,0 +1,867 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""The lenders beside the KV cache manager's automatic host tier: pages that move to host and back, +under their own request's suspension, other requests' pressure and a pool rebalance. Each check +first proves the pages did move, and runs once more against a lender breaking the rule it checks.""" + +import ctypes +from contextlib import contextmanager + +import numpy as np +import pytest +import torch + +from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import ( + StagingOptions, + attach_in_place, + attach_staging, +) + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="allocates KV cache pools") + +TPB = 32 +WINDOW = 64 +PROMPT = list(range(1000, 1097)) # three whole blocks and one token +END = 96 +BLOCKS = END // TPB +BYSTANDERS = [list(range(3000 + 200 * i, 3097 + 200 * i)) for i in range(2)] +WINDOWED_PROMPT = list(range(2000, 2161)) # five whole blocks and one token +WINDOWED_END = 160 +SOURCE, TARGET, SECOND = 1, 2, 3 +SCOPE = b"host-tier-suite" +# Only an assertion counts as a catch: a timeout or any other error fails the liar test. +CAUGHT = (AssertionError,) + + +def attach(mgr, *, fetch_tokens=END, max_fetches=1): + return attach_staging(mgr, scope=SCOPE, staging=StagingOptions(fetch_tokens, max_fetches)) + + +def attach_breaking(rules): + """An attach installing a staging lender whose ``rules`` (method name -> function) replace the + real ones.""" + + def attach_rule_breaker(mgr, *, fetch_tokens=END, max_fetches=1): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + cls = type("RuleBreaker", (_lender.Staging,), dict(rules)) + options = StagingOptions(fetch_tokens, max_fetches) + return _lender._attach_staging(mgr, scope=SCOPE, staging=options, cls=cls) + + return attach_rule_breaker + + +def real(name): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + return getattr(_lender.Staging, name) + + +def ready_view(lease, mgr, tries=3): + for _ in range(tries): + mgr._stream.synchronize() + view = lease.poll() + if view is not None: + return view + raise AssertionError(f"the lease never became ready: {lease.failure}") + + +def staged(lender, view): + """Per run, the bytes of every row's staging slot.""" + out = [] + for run in view.runs: + length = lender.parts[run.part].slot_bytes + out.append([ctypes.string_at(a, length) for a in run.addresses.tolist()]) + return out + + +def names(view): + return [[name.tobytes() for name in run.names] for run in view.runs] + + +@contextmanager +def to_host_and_back(kit, mgr, request, dev=None): + """Suspend ``request``; one-block requests take every GPU page, its pages go to host; free the + others but those on its old pages, resume it. Yields those others and the old pages. With + ``dev`` the others fill their pages on the stream without waiting for it.""" + kv = kit.kv(mgr, request) + old = [p for p in kit.pages(kv, 0) if p >= 0] + host_before = kit.tier_used(mgr, 1) + others = kit.Requests(mgr, dev) + try: + mgr.suspend_request(request) + assert others.allocate(kit.pool_pages(mgr), fill=dev is None, chunk=1) + assert set(old) <= others.pages(), "the request's pages stayed on the GPU" + assert kit.tier_used(mgr, 1) >= host_before + len(old), "its pages did not reach host" + others.free_all_but(set(old)) + assert mgr.resume_request(request) + now = [p for p in kit.pages(kv, 0) if p >= 0] + assert len(now) == len(old) and not set(now) & set(old), "it came back to its old pages" + yield others, old + finally: + others.free() + + +# -- (a) a fetch and a later publish across the request's own trip to host --------------------- + + +def check_a_fetch_stays_usable_across_a_trip_to_host(kit, host_tier_manager, attach): + with host_tier_manager() as mgr_a, host_tier_manager() as mgr_b: + source = kit.published(mgr_a, SOURCE, PROMPT) + lender_a = attach(mgr_a) + publish = lender_a.lend_read(source, 0, END) + publish_view = ready_view(publish, mgr_a) + target = kit.admitted(mgr_b, TARGET, PROMPT) + lender_b = attach(mgr_b) + lease = lender_b.lend_write(target, 0, END) + view = lease.poll() + assert view is not None + kit.fill_sentinel(mgr_b, target) + masks = kit.relay(lender_a, publish_view, lender_b, view) + dev = kit.DevicePages(mgr_b) + kit.warm_tier_moves(mgr_b) + with kit.held_stream(mgr_b._stream, strict=False) as gate: + # The copy into the lent pages waits behind the gate; the suspension and the move to + # host come after it on the stream, so the move carries the fetched bytes. + lease.mark_arrived(masks) + lease.release() + with to_host_and_back(kit, mgr_b, target, dev) as (others, old): + assert lender_b.readiness(target) is None, "settled before its copy ran" + gate.open() + mgr_b._stream.synchronize() + assert lender_b.readiness(target) == (END, 0), "not usable once its copy ran" + for ordinal in range(BLOCKS): + got = kit.digest(kit.page(mgr_b, target, 0, ordinal)) + want = kit.digest(kit.page(mgr_a, source, 0, ordinal)) + assert got == want, "the fetched bytes are lost" + sentinel = bytes([kit.SENTINEL]) * dev.page_bytes(0) + assert all(dev.read(0, slot) == sentinel for slot in old) + publish.release() + + +def test_a_fetch_stays_usable_across_its_request_s_trip_to_host(kit, host_tier_manager): + check_a_fetch_stays_usable_across_a_trip_to_host(kit, host_tier_manager, attach) + + +def test_the_check_catches_a_lender_tying_a_fetch_to_where_its_pages_were(kit, host_tier_manager): + def deliver(self, request_id, fetch, rows, copied): # also remembers the slots the rows had + real("_deliver")(self, request_id, fetch, rows, copied) + self.lent_rows = rows + + def where_they_were(self, manager, delivered, committed): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _manager + + rows = self.lent_rows + for lg, ordinals, slots in zip(rows.layer_groups, rows.ordinals, rows.device_slots): + moved = ordinals[_manager.pages(delivered.kv, lg)[ordinals] != slots] + delivered.blocks[lg][moved[moved < len(delivered.blocks[lg])]] = False + return real("_usable_until")(self, manager, delivered, committed) + + liar = attach_breaking({"_deliver": deliver, "_usable_until": where_they_were}) + with pytest.raises(CAUGHT, match="not usable once its copy ran"): + check_a_fetch_stays_usable_across_a_trip_to_host(kit, host_tier_manager, liar) + + +def check_a_publish_after_a_trip_to_host_reads_the_pages_where_they_are( + kit, host_tier_manager, attach +): + with host_tier_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + original = [kit.page(mgr, source, 0, o) for o in range(BLOCKS)] + lender = attach(mgr) + first = lender.lend_read(source, 0, END) + assert kit.digest(staged(lender, ready_view(first, mgr))) == kit.digest([original]) + first.release() + with to_host_and_back(kit, mgr, source): + again = lender.lend_read(source, 0, END) + got = kit.digest(staged(lender, ready_view(again, mgr))) + assert got == kit.digest([original]), "read a stale page" + again.release() + + +def test_a_publish_after_a_trip_to_host_reads_the_pages_where_they_are(kit, host_tier_manager): + check_a_publish_after_a_trip_to_host_reads_the_pages_where_they_are( + kit, host_tier_manager, attach + ) + + +def test_the_check_catches_a_lender_keeping_page_indices_across_lends(kit, host_tier_manager): + seen = {} + + def remembered(self, manager, kv, start, end): + key = (id(kv), start, end) + if key not in seen: + seen[key] = real("_rows")(self, manager, kv, start, end) + return seen[key] + + with pytest.raises(CAUGHT, match="read a stale page"): + check_a_publish_after_a_trip_to_host_reads_the_pages_where_they_are( + kit, host_tier_manager, attach_breaking({"_rows": remembered}) + ) + + +# -- (b) other requests' pressure while a fetch is lent ----------------------------------------- + + +def check_pressure_while_a_fetch_is_lent(kit, host_tier_manager, attach, moment): + """``moment``: ``copy_queued``, the target stays active and its marks' copy waits on the stream + while other requests push committed blocks to host; ``before_marks``, the target itself goes to + host and back between the grant and the marks.""" + with host_tier_manager() as mgr_a, host_tier_manager() as mgr_b: + source = kit.published(mgr_a, SOURCE, PROMPT) + lender_a = attach(mgr_a) + publish = lender_a.lend_read(source, 0, END) + publish_view = ready_view(publish, mgr_a) + kit.warm_tier_moves(mgr_b) + for i, tokens in enumerate(BYSTANDERS): # committed blocks, then no request holds them + mgr_b.free_resources(kit.published(mgr_b, 10 + i, tokens)) + target = kit.admitted(mgr_b, TARGET, PROMPT) + lender_b = attach(mgr_b) + lease = lender_b.lend_write(target, 0, END) + view = lease.poll() + assert view is not None + kit.fill_sentinel(mgr_b, target) + masks = kit.relay(lender_a, publish_view, lender_b, view) + kv = kit.kv(mgr_b, target) + lent = kit.pages(kv, 0)[:BLOCKS] + dev = kit.DevicePages(mgr_b) + sentinel = bytes([kit.SENTINEL]) * dev.page_bytes(0) + if moment == "copy_queued": + others = kit.Requests(mgr_b, dev) + host_before = kit.tier_used(mgr_b, 1) + try: + with kit.held_stream(mgr_b._stream, strict=False) as gate: + lease.mark_arrived(masks) + lease.release() + assert others.allocate(kit.pool_pages(mgr_b) - BLOCKS, fill=False, chunk=1) + moved = kit.tier_used(mgr_b, 1) - host_before + assert moved >= len(BYSTANDERS) * BLOCKS, "no bystander block reached host" + + assert kit.pages(kv, 0)[:BLOCKS] == lent, "the lent pages moved" + assert lender_b.readiness(target) is None, "settled before its copy ran" + gate.open() + assert lender_b.readiness(target) == (END, 0) + assert all(dev.read(0, slot) == sentinel for slot in others.pages()) + finally: + others.free() + else: + with to_host_and_back(kit, mgr_b, target) as (others, old): + assert old == lent + lease.mark_arrived(masks) + lease.release() + mgr_b._stream.synchronize() + assert all(dev.read(0, slot) == sentinel for slot in old), ( + "the marks' copy reached another request's page" + ) + assert lender_b.readiness(target) == (0, 0), "counts rows that never landed" + for ordinal in range(BLOCKS): + got = kit.digest(kit.page(mgr_b, target, 0, ordinal)) + expected = kit.page(mgr_a, source, 0, ordinal) if moment == "copy_queued" else sentinel + assert got == kit.digest(expected), "a page does not hold what its fetch says" + publish.release() + + +@pytest.mark.parametrize("moment", ["copy_queued", "before_marks"]) +def test_other_requests_pushing_blocks_to_host_leave_a_fetch_where_its_readiness_says( + kit, host_tier_manager, moment +): + check_pressure_while_a_fetch_is_lent(kit, host_tier_manager, attach, moment) + + +def test_the_check_catches_a_lender_settling_a_fetch_before_its_copy_ran(kit, host_tier_manager): + def at_the_marks(self, fetch): # settled once marked, whether or not the copy has run + return fetch.delivered + + with pytest.raises(CAUGHT, match="settled before its copy ran"): + check_pressure_while_a_fetch_is_lent( + kit, + host_tier_manager, + attach_breaking({"_report_settled": at_the_marks}), + "copy_queued", + ) + + +def test_the_check_catches_a_lender_trusting_the_cache_but_not_its_pages(kit, host_tier_manager): + def same_cache(self, kv, lease): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _manager + + alive = kv is not None and kv is lease._kv and _manager.cache_state(kv).active + return [np.full(len(o), alive, dtype=bool) for o in lease._rows.ordinals] + + with pytest.raises(CAUGHT, match="the marks' copy reached another request's page"): + check_pressure_while_a_fetch_is_lent( + kit, host_tier_manager, attach_breaking({"_still_lent": same_cache}), "before_marks" + ) + + +# -- (c) a cache whose pages are on host fails every lend at the call --------------------------- + +SERVED = "served bytes that are not its blocks'" +REACHED = "reached another request's page" + + +def use_as_a_backend(kit, mgr, lender, lease, reading, kind, kv, original): + """What a backend does with a lease that became ready: a read's rows are compared with the + blocks' bytes, a write's rows are filled. In place, rows are found in the manager's page table.""" + mgr._stream.synchronize() + view = lease.poll() + if view is None: + return + if kind == "staging": + if reading: + (run,) = view.runs + want = [[original[o] for o in run.ordinals.tolist()]] + assert kit.digest(staged(lender, view)) == kit.digest(want), SERVED + else: + kit.stage(lender, view, 0x5A) + lease.mark_arrived(view.row_masks(True)) + return + dev = kit.DevicePages(mgr) + for run in view.runs: + lg, slots = run.layer_group, kit.pages(kv, run.layer_group) + for ordinal in run.ordinals.tolist(): + if reading: + got = dev.read(lg, slots[ordinal]) + assert kit.digest(got) == kit.digest(original[ordinal]), SERVED + else: + dev.write(lg, slots[ordinal], bytes([0x5A]) * dev.page_bytes(lg)) + + +def check_a_cache_on_host_fails_lends_at_the_call(kit, host_tier_manager, kind, to_host=True): + """Both caches go to host before the lends. ``to_host=False`` leaves their pages on the GPU, + for a lender's rule breach to show without the host tier.""" + with host_tier_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + original = [kit.page(mgr, source, 0, o) for o in range(BLOCKS)] + target = kit.admitted(mgr, TARGET, BYSTANDERS[0]) + assert mgr.resize_context(target, target.context_remaining_length) + lender = attach(mgr, max_fetches=2) if kind == "staging" else attach_in_place(mgr) + writer = target if kind == "staging" else source + kv, kv_writer = kit.kv(mgr, source), kit.kv(mgr, writer) + old = [p for p in kit.pages(kv, 0) if p >= 0] + old_writer = [p for p in kit.pages(kv_writer, 0) if p >= 0] + host_before = kit.tier_used(mgr, 1) + others = kit.Requests(mgr) + try: + mgr.suspend_request(target) + mgr.suspend_request(source) + if to_host: + assert others.allocate(kit.pool_pages(mgr), chunk=1) + assert set(old) | set(old_writer) <= others.pages() + moved = len(set(old) | set(old_writer)) + assert kit.tier_used(mgr, 1) >= host_before + moved, "its pages did not reach host" + dev = kit.DevicePages(mgr) + theirs = {slot: dev.read(0, slot) for slot in others.pages()} + failed = [lender.lend_read(source, 0, END), lender.lend_write(writer, 0, END)] + for lease, reading, lent in zip(failed, (True, False), (kv, kv_writer)): + use_as_a_backend(kit, mgr, lender, lease, reading, kind, lent, original) + mgr._stream.synchronize() + now = {slot: dev.read(0, slot) for slot in theirs} + assert kit.digest(now) == kit.digest(theirs), REACHED + for lease in failed: + assert lease.failure is not None, "not failed at the call" + assert lease.poll() is None + if kind == "staging": + groups = kit.pool_group_ids(mgr) + assert [lender._free_slots(g) for g in groups] == [p.slots for p in lender.parts] + for lease in failed: + lease.release() + others.free_all_but(set(old)) + assert mgr.resume_request(source) and mgr.resume_request(target) + now = [p for p in kit.pages(kv, 0) if p >= 0] + assert not set(now) & set(old), "it came back to its old pages" + read = lender.lend_read(source, 0, END) + view = ready_view(read, mgr) + if kind == "staging": + assert kit.digest(staged(lender, view)) == kit.digest([original]) + else: + (run,) = view.runs + assert run.ordinals.tolist() == list(range(BLOCKS)) + write = lender.lend_write(writer, 0, END) + write_view = ready_view(write, mgr) + write.mark_arrived(write_view.row_masks()) + write.release() + read.release() + finally: + others.free() + + +@pytest.mark.parametrize("kind", ["staging", "in_place"]) +def test_a_cache_on_host_fails_every_lend_at_the_call_and_lends_again_once_resumed( + kit, host_tier_manager, kind +): + check_a_cache_on_host_fails_lends_at_the_call(kit, host_tier_manager, kind) + + +@pytest.mark.parametrize("pages", ["on_host", "on_gpu"]) +@pytest.mark.parametrize("kind", ["staging", "in_place"]) +def test_the_check_catches_a_lender_lending_a_suspended_cache_by_the_pages_it_lists( + kit, host_tier_manager, monkeypatch, kind, pages +): + """The lender reads a suspended cache as active and, in place, lends every page it lists. On + host those pages are others' and the contents catch it; on the GPU only the call's rule does.""" + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _manager + + state = _manager.cache_state + monkeypatch.setattr(_manager, "cache_state", lambda kv: state(kv)._replace(active=True)) + if kind == "in_place": + monkeypatch.setattr(_manager, "locked_pages", _manager.pages) + caught = f"{SERVED}|{REACHED}" if pages == "on_host" else "not failed at the call" + with pytest.raises(CAUGHT, match=caught): + check_a_cache_on_host_fails_lends_at_the_call( + kit, host_tier_manager, kind, to_host=pages == "on_host" + ) + + +# -- (d) a committed block's name does not depend on its tier ----------------------------------- + + +def check_a_block_keeps_its_name_on_either_tier(kit, host_tier_manager, attach): + with host_tier_manager() as mgr: + first = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr) + lease = lender.lend_read(first, 0, END) + view = ready_view(lease, mgr) + first_names, first_bytes = names(view), staged(lender, view) + lease.release() + old = kit.pages(kit.kv(mgr, first), 0)[:BLOCKS] + mgr.free_resources(first) # its whole blocks stay committed in the reuse tree + host_before = kit.tier_used(mgr, 1) + others = kit.Requests(mgr) + try: + assert others.allocate(kit.pool_pages(mgr), chunk=1) + assert set(old) <= others.pages() + assert kit.tier_used(mgr, 1) >= host_before + BLOCKS, "the blocks did not reach host" + others.free_all_but(set(old)) + second = kit.make_request(SECOND, PROMPT) + assert mgr.prepare_context(second) + kv = kit.kv(mgr, second) + assert kv.num_committed_tokens == END, "the blocks on host were not reused" + now = kit.pages(kv, 0)[:BLOCKS] + assert not set(now) & set(old), "the blocks came back to their old pages" + again = lender.lend_read(second, 0, END) + view = ready_view(again, mgr) + assert names(view) == first_names, "a block's name changed with its tier" + assert kit.digest(staged(lender, view)) == kit.digest(first_bytes) + again.release() + finally: + others.free() + + +def test_a_committed_block_has_one_name_on_the_gpu_and_after_a_trip_to_host(kit, host_tier_manager): + check_a_block_keeps_its_name_on_either_tier(kit, host_tier_manager, attach) + + +def test_the_check_catches_a_lender_naming_a_block_by_its_slot(kit, host_tier_manager): + def by_slot(self, rows, keys): + slot_keys = [ + np.repeat(slots.astype("= 0 for p in held), "the blocks behind the window hold no pages here" + lender = attach_in_place(mgr) + host_before = kit.tier_used(mgr, 1) + # Still on the GPU, before any pressure, a held page is left out all the same. + early = lender.lend_read(request, 0, 2 * TPB) + assert kit.pages(kv, sliding)[:end] == held and kit.tier_used(mgr, 1) == host_before + runs = {run.layer_group: run.ordinals.tolist() for run in early.poll().runs} + assert runs == {1 - sliding: [0, 1], sliding: []}, "lent a held page still on the GPU" + early.release() + others = kit.Requests(mgr) + try: + while others.allocate(1, chunk=1): + pass + gone = [o for o in range(end) if kit.pages(kv, sliding)[o] != held[o]] + assert set(gone) >= {0, 1}, "the held pages stayed on the GPU" + assert kit.tier_used(mgr, 1) >= host_before + len(gone) + read = lender.lend_read(request, 0, 2 * TPB) # a history of 64 reads blocks 0 and 1 + runs = {run.layer_group: run.ordinals.tolist() for run in read.poll().runs} + assert runs == {1 - sliding: [0, 1], sliding: []}, "lent a page it only holds" + read.release() + write = lender.lend_write(request, 0, 2 * TPB) + assert write.failure is not None, "a write into pages only held went through" + write.release() + finally: + others.free() + + +def test_in_place_below_the_history_lends_only_pages_its_cache_locks(kit, host_tier_manager): + check_in_place_below_the_history_lends_only_locked_pages(kit, host_tier_manager) + + +def test_the_check_catches_a_lender_lending_held_pages_below_the_history( + kit, host_tier_manager, monkeypatch +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender, _manager + + locked = _manager.locked_pages + + def every_page(kv, layer_group): # every block's page, held ones included + if not kv.is_active: + return locked(kv, layer_group) + return _manager.pages(kv, layer_group) + + monkeypatch.setattr(_lender._manager, "locked_pages", every_page) + with pytest.raises(CAUGHT, match="lent a held page still on the GPU"): + check_in_place_below_the_history_lends_only_locked_pages(kit, host_tier_manager) + + +def test_the_check_catches_a_lender_lending_held_pages_while_they_stay_on_the_gpu( + kit, host_tier_manager, monkeypatch +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender, _manager + + locked = _manager.locked_pages + + def while_nothing_is_on_host(kv, layer_group): # a held page taken as safe on the GPU + stats = kv.manager.get_storage_statistics(1) + if any(int(s.total) != int(s.free) for s in stats): + return locked(kv, layer_group) + return _manager.pages(kv, layer_group) + + monkeypatch.setattr(_lender._manager, "locked_pages", while_nothing_is_on_host) + with pytest.raises(CAUGHT, match="lent a held page still on the GPU"): + check_in_place_below_the_history_lends_only_locked_pages(kit, host_tier_manager) + + +def test_the_check_catches_a_lender_lending_held_pages_once_some_are_on_host( + kit, host_tier_manager, monkeypatch +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender, _manager + + locked = _manager.locked_pages + + def once_on_host(kv, layer_group): # right before any pressure, wrong after it + stats = kv.manager.get_storage_statistics(1) + if any(int(s.total) != int(s.free) for s in stats): + return _manager.pages(kv, layer_group) + return locked(kv, layer_group) + + monkeypatch.setattr(_lender._manager, "locked_pages", once_on_host) + with pytest.raises(CAUGHT, match="lent a page it only holds"): + check_in_place_below_the_history_lends_only_locked_pages(kit, host_tier_manager) + + +# -- (e) a pool rebalance while a fetch is lent ------------------------------------------------- + + +class DeviceCopyGuard: + """The CUDA driver, except that a memcpy into device memory outside ``allowed()`` is recorded and + not run: after a pool shrinks, the slot it names may be unmapped.""" + + def __init__(self, real_driver, device_ranges, allowed): + self._real = real_driver + self._device = device_ranges # every byte a device pool ever held + self._allowed = allowed # () -> [(begin, end)] the copies may write + self.into_device = 0 + self.stray = [] + + @staticmethod + def _covered(begin, end, ranges): + for lo, hi in sorted(ranges): + if lo <= begin < hi: + begin = hi + if begin >= end: + return True + return False + + def __getattr__(self, name): + found = getattr(self._real, name) + if name != "cuMemcpyAsync": + return found + + def call(dst, src, nbytes, stream): + begin = int(dst) + if any(lo <= begin < hi for lo, hi in self._device): + self.into_device += 1 + if not self._covered(begin, begin + int(nbytes), self._allowed()): + self.stray.append((begin, int(nbytes))) + return (self._real.CUresult.CUDA_SUCCESS,) + return found(dst, src, nbytes, stream) + + return call + + +def page_ranges(kit, mgr, kv, layer_groups): + """Byte ranges of the pages ``kv`` locks in ``layer_groups``, in every pool of their groups.""" + group_of = {int(pg.pool_group_index): pg for pg in mgr.impl.pool_group_descs} + pool_group = kit.pool_group_of(mgr) + out = [] + for lg in layer_groups: + pg = group_of[pool_group[lg]] + slots = [int(s) for s in kv.get_base_page_indices(lg)[: kv.num_blocks] if int(s) >= 0] + for pool in pg.pools: + base, width = int(pool.base_address), int(pool.slot_bytes) + out.extend((base + s * width, base + (s + 1) * width) for s in slots) + return out + + +def check_a_rebalance_under_a_fetch_moves_no_copy_off_the_pages_still_lent( + kit, host_tier_manager, monkeypatch, attach +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + from tensorrt_llm.runtime.kv_cache_manager_v2 import _introspection + + shape = dict(windows=[WINDOW, 256], num_layers=3, max_tokens=4 * kit.POOL_TOKENS) + with host_tier_manager(**shape) as mgr_a, host_tier_manager(**shape) as mgr_b: + source = kit.published(mgr_a, SOURCE, WINDOWED_PROMPT) + lender_a = attach(mgr_a, fetch_tokens=WINDOWED_END) + publish = lender_a.lend_read(source, 0, WINDOWED_END) + publish_view = ready_view(publish, mgr_a) + full = kit.windows(mgr_b).index(None) + pg_full = kit.pool_group_of(mgr_b)[full] + device = [] + for pg in mgr_b.impl.pool_group_descs: + for pool in pg.pools: + base = int(pool.base_address) + device.append((base, base + int(pg.num_slots) * int(pool.slot_bytes))) + # Fillers take the low slots, so the target's full-attention pages sit high in their pool. + fillers = kit.Requests(mgr_b) + stats = mgr_b.impl.get_storage_statistics + while stats(0)[pg_full].total - stats(0)[pg_full].free < 72 and fillers.allocate(7): + pass + target = kit.admitted(mgr_b, TARGET, WINDOWED_PROMPT) + lender_b = attach(mgr_b, fetch_tokens=WINDOWED_END) + lease = lender_b.lend_write(target, 0, WINDOWED_END) + view = lease.poll() + assert view is not None + kit.fill_sentinel(mgr_b, target) + masks = kit.relay(lender_a, publish_view, lender_b, view) + kv = kit.kv(mgr_b, target) + lent = [kit.pages(kv, lg) for lg in range(2)] + fillers.free() + # The executor's rebalance: suspend every active request, adjust the pools, resume them. + mgr_b.suspend_request(target) + _introspection.force_rebalance_precondition(mgr_b.impl, skew=0.05 if pg_full == 0 else 20) + assert mgr_b.impl.need_adjustment + mgr_b.impl.adjust() + assert mgr_b.resume_request(target) + now = [kit.pages(kv, lg) for lg in range(2)] + slots_now = { + int(pg.pool_group_index): int(pg.num_slots) for pg in mgr_b.impl.pool_group_descs + } + gone = [p for p in lent[full] if p >= slots_now[pg_full]] + assert now[full] != lent[full] and gone, "the rebalance left the lent pages where they were" + guard = DeviceCopyGuard(_lender.drv, device, lambda: page_ranges(kit, mgr_b, kv, range(2))) + monkeypatch.setattr(_lender, "drv", guard) + lease.mark_arrived(masks) + lease.release() + mgr_b._stream.synchronize() + assert guard.stray == [], "the marks' copy wrote outside the pages the target locks" + assert guard.into_device, "nothing was copied: the check saw no copy" + assert lender_b.readiness(target) == (0, WINDOWED_END), "counts rows that never landed" + publish.release() + + +def test_a_rebalance_under_a_fetch_moves_no_copy_off_the_pages_still_lent( + kit, host_tier_manager, monkeypatch +): + check_a_rebalance_under_a_fetch_moves_no_copy_off_the_pages_still_lent( + kit, host_tier_manager, monkeypatch, attach + ) + + +def test_the_check_catches_a_lender_copying_into_pages_a_rebalance_moved( + kit, host_tier_manager, monkeypatch +): + def same_cache(self, kv, lease): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _manager + + alive = kv is not None and kv is lease._kv and _manager.cache_state(kv).active + return [np.full(len(o), alive, dtype=bool) for o in lease._rows.ordinals] + + with pytest.raises(CAUGHT, match="the marks' copy wrote outside the pages the target locks"): + check_a_rebalance_under_a_fetch_moves_no_copy_off_the_pages_still_lent( + kit, host_tier_manager, monkeypatch, attach_breaking({"_still_lent": same_cache}) + ) + + +# -- (f) a mark arriving after the window passed the lent blocks -------------------------------- + +LATE_PROMPT = list(range(2000, 2161)) # five whole blocks and one token +LATE_END = 2 * TPB # the lease covers blocks 0 and 1 +LATE_HISTORY = 5 * TPB # a served prefix moves the history here: blocks 0 to 2 leave the window +FETCHED = 0xAB + + +class CopyLog: + """The CUDA driver, recording every ``cuMemcpyAsync`` destination; every copy still runs.""" + + def __init__(self, real_driver): + self._real = real_driver + self.copies = [] + + def __getattr__(self, name): + found = getattr(self._real, name) + if name != "cuMemcpyAsync": + return found + + def call(dst, src, nbytes, stream): + self.copies.append(int(dst)) + return found(dst, src, nbytes, stream) + + return call + + +def device_slot(mgr, address): + """``(pool group, slot)`` of a device address, or ``None`` outside every device pool.""" + for pg in mgr.impl.pool_group_descs: + for pool in pg.pools: + base, width = int(pool.base_address), int(pool.slot_bytes) + if base <= address < base + int(pg.num_slots) * width: + return int(pg.pool_group_index), (address - base) // width + return None + + +def locked_by(kit, mgr, requests): + """``{(pool group, slot): request id}`` of every page ``requests`` lock, every layer group.""" + group_of = kit.pool_group_of(mgr) + out = {} + for request in requests: + kv = kit.kv(mgr, request) + for lg in range(kit.num_layer_groups(mgr)): + for slot in kv.get_base_page_indices(lg)[: kv.num_blocks]: + if int(slot) >= 0: + out[(group_of[lg], int(slot))] = request.py_request_id + return out + + +def late_mark(kit, host_tier_manager, monkeypatch, attach, windows, fillers, keep): + """One run, ``fillers`` one-block requests allocated first and kept or freed by ``keep``, which + moves the slots the target gets. What the marks' copy wrote, or ``None`` when no lent block went + to host under its lent GPU slot's number.""" + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + with host_tier_manager(windows=windows) as mgr: + sliding = kit.windows(mgr).index(WINDOW) + group = kit.pool_group_of(mgr)[sliding] + first, others = kit.Requests(mgr), kit.Requests(mgr) + others._next = 500 + try: + assert first.allocate(fillers, chunk=1) if fillers else True + if not keep: + first.free() + target = kit.admitted(mgr, TARGET, LATE_PROMPT) + lender = attach(mgr, fetch_tokens=LATE_END) + lease = lender.lend_write(target, 0, LATE_END) + view = lease.poll() + assert view is not None, lease.failure + kv = kit.kv(mgr, target) + lent = kit.pages(kv, sliding)[:2] + assert [int(s) for s in kv.get_base_page_indices(sliding)[:2]] == lent + kit.stage(lender, view, FETCHED) + # Before the marks, a served prefix moves the history past the lent blocks of the + # sliding group, as a KV connector's reservation does: they keep their pages only held. + assert mgr._resize_for_connector_prefix(target, kv, LATE_HISTORY, LATE_HISTORY + 1) + stale = kit.stale_blocks(mgr, sliding, kv.history_length) + assert stale[0] <= 0 and stale[1] >= 2, "the window did not pass the lent blocks" + assert all(int(s) < 0 for s in kv.get_base_page_indices(sliding)[:2]), "still locked" + assert kit.pages(kv, sliding)[:2] == lent, "the lent blocks do not hold their pages" + # Other requests take every GPU page left: the held pages go to host. + host_before = kit.tier_used(mgr, 1) + while others.allocate(4, chunk=4): + pass + while others.allocate(1, chunk=1): + pass + now = kit.pages(kv, sliding)[:2] + owners = locked_by(kit, mgr, others.held + first.held) + moved = [o for o in range(2) if (group, lent[o]) in owners and now[o] >= 0] + assert kit.tier_used(mgr, 1) >= host_before + len(moved), "nothing reached host" + aliased = [o for o in moved if now[o] == lent[o]] + if not aliased: + lease.mark_arrived(view.row_masks(False)) + lease.release() + return None + dev = kit.DevicePages(mgr) + watched = {lent[o]: owners[(group, lent[o])] for o in aliased} + before = {s: dev.read(sliding, s) for s in watched} + log = CopyLog(_lender.drv) + monkeypatch.setattr(_lender, "drv", log) + lease.mark_arrived(view.row_masks(True)) + lease.release() + mgr._stream.synchronize() + monkeypatch.undo() + written = [device_slot(mgr, dst) for dst in log.copies] + hit = sorted({w for w in written if w is not None and w in owners}) + changed = sorted(watched[s] for s in watched if dev.read(sliding, s) != before[s]) + return dict(hit=hit, changed=changed) + finally: + others.free() + first.free() + + +def check_a_late_mark_writes_no_slot_its_window_left(kit, host_tier_manager, monkeypatch, attach): + # Slots go out lowest first, so a short search finds a lent block on host under its lent GPU + # slot's number while another request locks that GPU slot. + runs = [(w, n, k) for w in ([WINDOW], [WINDOW, 256]) for k in (True, False) for n in range(4)] + for windows, fillers, keep in runs: + result = late_mark(kit, host_tier_manager, monkeypatch, attach, windows, fillers, keep) + if result is not None: + break + else: + raise RuntimeError("no run put a lent block on host under its lent GPU slot's number") + assert not result["hit"], f"the marks' copy wrote slots other requests lock: {result['hit']}" + assert not result["changed"], f"requests {result['changed']} now hold the fetch" + + +def test_a_late_mark_writes_no_slot_its_window_left(kit, host_tier_manager, monkeypatch): + check_a_late_mark_writes_no_slot_its_window_left(kit, host_tier_manager, monkeypatch, attach) + + +def test_the_check_catches_a_lender_matching_pages_by_slot_number_alone( + kit, host_tier_manager, monkeypatch +): + def same_number(self, kv, lease): # any page of the block, on any tier, under the lent number + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _manager + + rows = lease._rows + alive = kv is not None and kv is lease._kv and _manager.cache_state(kv).active + masks = [] + for lg, ordinals, slots in zip(rows.layer_groups, rows.ordinals, rows.device_slots): + same = np.zeros(len(ordinals), dtype=bool) + if alive: + pages = _manager.pages(kv, lg) + inside = ordinals < len(pages) + same[inside] = pages[ordinals[inside]] == slots[inside] + masks.append(same) + return masks + + with pytest.raises(CAUGHT, match="the marks' copy wrote slots other requests lock"): + check_a_late_mark_writes_no_slot_its_window_left( + kit, host_tier_manager, monkeypatch, attach_breaking({"_still_lent": same_number}) + ) diff --git a/tests/unittest/_torch/executor/kv_cache/sharing/test_in_place_lender.py b/tests/unittest/_torch/executor/kv_cache/sharing/test_in_place_lender.py new file mode 100644 index 000000000000..a7721bd46c5e --- /dev/null +++ b/tests/unittest/_torch/executor/kv_cache/sharing/test_in_place_lender.py @@ -0,0 +1,853 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""The in-place lender over real KV cache managers, through the public API: views of a request's +own pages, the two explicit failures, and loans kept across the request's free and the manager's +shutdown. Oracles read pages and the allocator directly; each has a negative control.""" + +import gc +import threading +import weakref +from contextlib import contextmanager + +import numpy as np +import pytest +import torch + +from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import ( + InPlaceLender, + Lease, + StagingLender, + StagingOptions, + attach_in_place, + attach_staging, +) + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="allocates KV cache pools") + +TPB = 32 +WINDOW = 64 +PROMPT = list(range(1000, 1097)) # three whole blocks and a one-token tail: four pages +END = len(PROMPT) +WINDOWED_PROMPT = list(range(2000, 2161)) # five whole blocks and a one-token tail +WINDOWED_END = len(WINDOWED_PROMPT) +SOURCE, TARGET = 1, 2 + + +def ceil_blocks(tokens: int) -> int: + return -(-tokens // TPB) + + +def check_view(kit, view, kv, expected): + """``expected``: layer group -> ordinals, every layer group listed. Rows carry no names, + addresses or part, and each is a block with a page of its own.""" + assert {run.layer_group: run.ordinals.tolist() for run in view.runs} == { + lg: list(ordinals) for lg, ordinals in expected.items() + } + for run in view.runs: + assert run.names is None and run.addresses is None and run.part is None + own = kit.pages(kv, run.layer_group) + assert all(own[o] >= 0 for o in run.ordinals.tolist()) + + +def in_turn(kit, call, name): + """``call`` on its own thread, joined; like the executor's threads, it used CUDA before.""" + + def run(): + torch.cuda.synchronize() + return call() + + return kit.on_thread(run, name) + + +@contextmanager +def lent_and_freed(kit, mgr, lender, write=False): + """``SOURCE`` computed ``PROMPT``, a lease lends all of it, and ``SOURCE`` is freed. Yields the + lease, the freed cache and its pages.""" + request = kit.published(mgr, SOURCE, PROMPT) + kv = kit.kv(mgr, request) + lent = set(kit.pages(kv, 0)) + assert len(lent) == ceil_blocks(END) and kit.pool_pages(mgr) - len(lent) >= 1 + lease = (lender.lend_write if write else lender.lend_read)(request, 0, END) + assert lease.poll() is not None + mgr.free_resources(request) + assert kit.kv(mgr, request) is None + yield lease, kv, lent + + +# -- what it lends ---------------------------------------------------------------------------- + + +def test_an_in_place_lender_lends_and_promises_nothing_more(real_manager): + with real_manager() as mgr: + lender = attach_in_place(mgr) + assert isinstance(lender, InPlaceLender) and not isinstance(lender, StagingLender) + assert not hasattr(lender, "readiness") and not hasattr(lender, "parts") + with pytest.raises(ValueError, match="already attached"): + attach_in_place(mgr) + with pytest.raises(ValueError, match="already attached"): + attach_staging(mgr, scope=b"scope", staging=StagingOptions(TPB)) + + +def test_a_read_is_ready_at_its_first_poll_without_waiting_for_the_stream(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + kv = kit.kv(mgr, source) + lender = attach_in_place(mgr) + with kit.held_stream(mgr._stream) as gate: + lease = lender.lend_read(source, 0, END) + assert isinstance(lease, Lease) + view = lease.poll() + assert view is not None, "an in-place lease waited for the stream" + assert lease.poll() is view + gate.open() + check_view(kit, view, kv, {0: range(ceil_blocks(END))}) + assert END % TPB and view.runs[0].ordinals[-1] == END // TPB, "the partial last block" + pieces = [ + (lender.lend_read(source, TPB, 2 * TPB), [1]), + (lender.lend_read(source, 2 * TPB, END), [2, 3]), + (lender.lend_read(source, 5, 40), [0, 1]), # any token range + ] + for piece, ordinals in pieces: + check_view(kit, piece.poll(), kv, {0: ordinals}) + for held in [lease] + [piece for piece, _ in pieces]: + held.release() + with pytest.raises(RuntimeError): + lease.poll() + + +def test_a_windowed_view_keeps_only_the_blocks_the_window_reads(kit, real_manager): + with real_manager(windows=[WINDOW, 256]) as mgr: + source = kit.published(mgr, SOURCE, WINDOWED_PROMPT) + kv = kit.kv(mgr, source) + lender = attach_in_place(mgr) + windows = kit.windows(mgr) + sliding, full = windows.index(WINDOW), windows.index(None) + beg, end = kit.stale_blocks(mgr, sliding, WINDOWED_END) + assert end > beg == 0, "the window must have left some blocks behind" + blocks = range(ceil_blocks(WINDOWED_END)) + lease = lender.lend_read(source, 0, WINDOWED_END) + in_window = [o for o in blocks if not beg <= o < end] + check_view(kit, lease.poll(), kv, {full: blocks, sliding: in_window}) + lease.release() + + +def test_a_generation_view_leaves_out_paged_blocks_the_window_no_longer_reads(kit, real_manager): + with real_manager(windows=[WINDOW, 256]) as mgr: + request = kit.make_request(TARGET, WINDOWED_PROMPT) + assert mgr.prepare_context(request) + assert mgr.resize_context(request, request.context_remaining_length) + kv = kit.kv(mgr, request) + lender = attach_in_place(mgr) + sliding = kit.windows(mgr).index(WINDOW) + beg, end = kit.stale_blocks(mgr, sliding, WINDOWED_END) + own = kit.pages(kv, sliding) + assert end > beg and all(own[o] >= 0 for o in range(beg, end)), ( + "the blocks below the window must still hold pages here" + ) + blocks = range(ceil_blocks(WINDOWED_END)) + lease = lender.lend_write(request, 0, WINDOWED_END) + in_window = [o for o in blocks if not beg <= o < end] + check_view(kit, lease.poll(), kv, {1 - sliding: blocks, sliding: in_window}) + lease.release() + + +def test_an_mla_view_lists_every_block_with_its_tail(kit, real_manager): + from tensorrt_llm.bindings import DataType + from tensorrt_llm.bindings.internal.batch_manager import CacheType + + mla = dict(kv_cache_type=CacheType.SELFKONLY, num_kv_heads=1, head_dim=128, dtype=DataType.BF16) + with real_manager(**mla) as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach_in_place(mgr) + lease = lender.lend_read(source, 0, END) + check_view(kit, lease.poll(), kit.kv(mgr, source), {0: range(ceil_blocks(END))}) + lease.release() + + +def test_a_write_changes_neither_the_cache_nor_its_pages(kit, real_manager): + with real_manager() as mgr: + kit.published(mgr, SOURCE, PROMPT) + request = kit.make_request(TARGET, PROMPT) + assert mgr.prepare_context(request) + kv = kit.kv(mgr, request) + assert kv.num_committed_tokens >= TPB, "the prefix must be reused" + assert mgr.resize_context(request, request.context_remaining_length) + lender = attach_in_place(mgr) + shape = (kv.capacity, kv.history_length, kv.num_committed_tokens) + lease = lender.lend_write(request, 0, END) # starts inside the reused prefix + view = lease.poll() + check_view(kit, view, kv, {0: range(ceil_blocks(END))}) + assert (kv.capacity, kv.history_length, kv.num_committed_tokens) == shape, "grown" + dev = kit.DevicePages(mgr) + slots = [kit.pages(kv, 0)[o] for o in view.runs[0].ordinals.tolist()] + content = [dev.read(0, slot) for slot in slots] + lease.mark_arrived(view.row_masks(True)) + mgr._stream.synchronize() + now = kit.digest([dev.read(0, slot) for slot in slots]) + assert now == kit.digest(content), "marks copied something" + assert (kv.capacity, kv.history_length, kv.num_committed_tokens) == shape + lease.release() + + +def test_a_view_may_end_past_the_committed_tokens_with_reuse_off(kit, real_manager): + from tensorrt_llm._torch.pyexecutor.scheduler import ScheduledRequests + + with real_manager(enable_block_reuse=False) as mgr: + request = kit.make_request(SOURCE, PROMPT) + assert mgr.prepare_context(request) + assert mgr.resize_context(request, request.context_remaining_length) + request.move_to_next_context_chunk() + batch = ScheduledRequests() + batch.append_context_request(request) + mgr.update_context_resources(batch) + kv = kit.kv(mgr, request) + assert kv.num_committed_tokens < END + lender = attach_in_place(mgr) + lease = lender.lend_read(request, 0, END) + check_view(kit, lease.poll(), kv, {0: range(ceil_blocks(END))}) + lease.release() + + +# -- failures --------------------------------------------------------------------------------- + + +def test_blocks_without_a_page_are_left_out_of_a_read_and_fail_a_write(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + kv = kit.kv(mgr, source) + blocks = ceil_blocks(END) + assert len(kit.pages(kv, 0)) == blocks + lender = attach_in_place(mgr) + beyond = (blocks + 1) * TPB + read = lender.lend_read(source, 0, beyond) + check_view(kit, read.poll(), kv, {0: range(blocks)}) + empty = lender.lend_read(source, blocks * TPB, beyond) + assert empty.poll().num_rows == 0 + write = lender.lend_write(source, 0, beyond) + assert write.failure is not None and write.poll() is None + whole = lender.lend_write(source, 0, blocks * TPB) # every block it touches has a page + assert whole.poll() is not None + for lease in (read, empty, whole): + lease.release() + mgr.free_resources(source) + assert kit.closed(kv), "a failed lease holds no loan" + write.release() + + +def test_a_suspended_cache_fails_both_lends_at_the_call(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach_in_place(mgr) + mgr.suspend_request(source) + assert not kit.kv(mgr, source).is_active + for lend in (lender.lend_read, lender.lend_write): + lease = lend(source, 0, END) + assert lease.failure is not None and lease.poll() is None + lease.release() + assert mgr.resume_request(source) # active again, both go through + for lend in (lender.lend_read, lender.lend_write): + lease = lend(source, 0, END) + assert lease.failure is None and lease.poll() is not None + lease.release() + + +def test_only_a_negative_or_reversed_range_is_an_argument_error(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach_in_place(mgr) + for lend in (lender.lend_read, lender.lend_write): + for start, end in ((-1, TPB), (TPB, 0), (0, -1)): + with pytest.raises(ValueError): + lend(source, start, end) + nobody = lend(kit.make_request(9, PROMPT), 0, TPB) + assert nobody.failure is not None and nobody.poll() is None + nobody.release() + mgr.shutdown() + for lend in (lender.lend_read, lender.lend_write): + lease = lend(source, 0, END) + assert lease.failure is not None and lease.poll() is None + lease.release() + + +def test_mark_arrived_checks_the_shapes_and_nothing_else(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach_in_place(mgr) + read = lender.lend_read(source, 0, END) + read_view = read.poll() + with pytest.raises(RuntimeError): + read.mark_arrived(read_view.row_masks()) + write = lender.lend_write(source, 0, END) + with pytest.raises(RuntimeError): + write.mark_arrived(()) # before poll() gave the view + view = write.poll() + for masks in ((), (np.ones(len(view.runs[0]) + 1, bool),)): + with pytest.raises(ValueError): + write.mark_arrived(masks) + write.release() + write.mark_arrived(view.row_masks(True)) # after the release too + with pytest.raises(RuntimeError): + write.mark_arrived(view.row_masks(True)) + read.release() + + +# -- the request's free ----------------------------------------------------------------------- + + +def check_a_freed_request_keeps_its_lent_pages_until_the_release(kit, real_manager): + with real_manager(max_tokens=kit.POOL_TOKENS) as mgr: + dev = kit.DevicePages(mgr) + lender = attach_in_place(mgr) + with lent_and_freed(kit, mgr, lender) as (lease, kv, lent): + content = {slot: dev.read(0, slot) for slot in lent} + assert not kit.closed(kv), "a freed cache stays open while lent" + assert kit.taken_by_others(mgr, lent) == set(), "another request got a lent page" + now = kit.digest({slot: dev.read(0, slot) for slot in lent}) + assert now == kit.digest(content), "a lent page changed" + lease.release() + assert kit.closed(kv), "the last release closes the cache in its call" + assert kit.whole_pool_goes_to_others(mgr, lent), "the freed pages are reused" + + +def test_a_freed_request_keeps_its_lent_pages_until_the_release(kit, real_manager): + check_a_freed_request_keeps_its_lent_pages_until_the_release(kit, real_manager) + + +def test_the_check_catches_a_lender_letting_the_free_close_a_lent_cache( + kit, real_manager, monkeypatch +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + monkeypatch.setattr(_lender.InPlace, "_on_free", lambda self, rid, kv, after: False) + with pytest.raises(AssertionError, match="a freed cache stays open while lent"): + check_a_freed_request_keeps_its_lent_pages_until_the_release(kit, real_manager) + + +def test_the_page_oracle_sees_a_page_the_free_returned(kit, real_manager, monkeypatch): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + monkeypatch.setattr(_lender.InPlace, "_on_free", lambda self, rid, kv, after: False) + with real_manager(max_tokens=kit.POOL_TOKENS) as mgr: + lender = attach_in_place(mgr) + with lent_and_freed(kit, mgr, lender) as (lease, kv, lent): + assert kit.closed(kv) + assert kit.taken_by_others(mgr, lent), "the oracle cannot see a returned page" + lease.release() + + +@pytest.mark.parametrize("mark_first", [True, False], ids=["mark_first", "release_first"]) +def test_marks_and_the_release_may_come_in_either_order(kit, real_manager, mark_first): + with real_manager(max_tokens=kit.POOL_TOKENS) as mgr: + lender = attach_in_place(mgr) + with lent_and_freed(kit, mgr, lender, write=True) as (lease, kv, lent): + view = lease.poll() + if mark_first: + lease.mark_arrived(view.row_masks(True)) + assert not kit.closed(kv), "marks do not end the loan" + assert kit.taken_by_others(mgr, lent) == set() + lease.release() + assert kit.closed(kv) + else: + lease.release() + assert kit.closed(kv), "the release ends the loan" + lease.mark_arrived(view.row_masks(True)) + assert kit.closed(kv) + assert kit.whole_pool_goes_to_others(mgr, lent) + + +def test_the_pages_stay_until_the_last_of_several_leases_ends(kit, real_manager): + with real_manager(max_tokens=kit.POOL_TOKENS) as mgr: + request = kit.published(mgr, SOURCE, PROMPT) + kv = kit.kv(mgr, request) + lent = set(kit.pages(kv, 0)) + lender = attach_in_place(mgr) + # Never polled: a loan opens at the call. + write = lender.lend_write(request, 0, END) + read = lender.lend_read(request, 0, 2 * TPB) + mgr.free_resources(request) + write.release() + assert not kit.closed(kv) and kit.taken_by_others(mgr, lent) == set() + read.release() + assert kit.closed(kv) + + +def next_owner_s_row_across_the_kept_close(kit, mgr): + """``SOURCE`` is lent, freed, and its index slot taken by another request; then the lease ends + and the kept cache closes. The new owner's host page table row before and after.""" + request = kit.published(mgr, SOURCE, PROMPT) + kv = kit.kv(mgr, request) + index = mgr.index_mapper.get_index(SOURCE) + lender = attach_in_place(mgr) + lease = lender.lend_read(request, 0, END) + mgr.free_resources(request) + others = kit.Requests(mgr) + assert others.allocate(3) + (other,) = others.held + assert mgr.index_mapper.get_index(other.py_request_id) == index, "the slot is reused" + row = mgr.host_kv_cache_block_offsets[0, index * mgr.max_beam_width] + before = row.clone() + lease.release() + assert kit.closed(kv) + after = row.clone() + others.free() + return before, after + + +def test_a_lent_request_gives_up_its_index_slot_at_its_free(kit, real_manager): + with real_manager(max_tokens=kit.POOL_TOKENS) as mgr: + before, after = next_owner_s_row_across_the_kept_close(kit, mgr) + assert torch.equal(after, before), "closing the kept cache wrote into the next owner's row" + + +def test_the_row_oracle_sees_a_close_writing_into_the_next_owner(kit, real_manager, monkeypatch): + from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import KVCacheManagerV2 + + def free_lent_attached(self, request_id, kv_cache): + if request_id in self._early_freed_index_requests: + self._early_freed_index_requests.discard(request_id) + return + self.index_mapper.remove_sequence(request_id) + + monkeypatch.setattr(KVCacheManagerV2, "_free_lent", free_lent_attached) + with real_manager(max_tokens=kit.POOL_TOKENS) as mgr: + before, after = next_owner_s_row_across_the_kept_close(kit, mgr) + assert not torch.equal(after, before), "the oracle cannot see a write into the row" + + +def test_a_kept_cache_closes_before_its_stats_exclusion_is_cleared(kit, real_manager): + with real_manager(max_tokens=kit.POOL_TOKENS) as mgr: + request = kit.published(mgr, SOURCE, PROMPT) + kv = kit.kv(mgr, request) + mgr.impl.mark_stats_excluded(SOURCE) + lender = attach_in_place(mgr) + lease = lender.lend_read(request, 0, END) + impl = mgr.impl + spy = mgr.impl = kit.StatsSpy(impl, kv) + try: + mgr.free_resources(request) + assert spy.cleared == [] and impl.is_stats_excluded(SOURCE) + lease.release() + assert spy.cleared == [(SOURCE, True)], "cleared before the close" + assert not impl.is_stats_excluded(SOURCE) + finally: + mgr.impl = impl + + +# -- the manager's shutdown ------------------------------------------------------------------- + + +@pytest.mark.parametrize("freed", [True, False], ids=["request_freed", "request_live"]) +def test_shutdown_keeps_what_a_lease_still_lends_until_exit(kit, real_manager, monkeypatch, freed): + with real_manager(max_tokens=kit.POOL_TOKENS) as mgr: + request = kit.published(mgr, SOURCE, PROMPT) + kv = kit.kv(mgr, request) + dev = kit.DevicePages(mgr) + lender = attach_in_place(mgr) + lease = lender.lend_read(request, 0, END) + slots = [kit.pages(kv, 0)[o] for o in lease.poll().runs[0].ordinals.tolist()] + content = [dev.read(0, slot) for slot in slots] + if freed: + mgr.free_resources(request) + spy = mgr.impl = kit.ShutdownSpy(mgr.impl) + warnings = kit.lender_warnings(monkeypatch) + mgr.shutdown() + assert not spy.shut_down, "the pools holding a lent page were destroyed" + assert warnings, "keeping caches until exit is logged" + kept = kit.retained() + assert any(o is spy for o in kept) and any(o is kv for o in kept) + assert not kit.closed(kv) + now = kit.digest([dev.read(0, slot) for slot in slots]) + assert now == kit.digest(content), "the lent bytes stay readable" + lease.release() + assert not kit.closed(kv), "a release after the shutdown closes nothing" + mgr.shutdown() + assert not spy.shut_down, "a later shutdown keeps them too" + + +@pytest.mark.parametrize("leases", [0, 2], ids=["never_lent", "every_lease_ended"]) +def test_without_an_open_loan_free_and_shutdown_are_as_without_a_lender( + kit, real_manager, monkeypatch, leases +): + with real_manager(max_tokens=kit.POOL_TOKENS) as mgr: + request = kit.published(mgr, SOURCE, PROMPT) + kv = kit.kv(mgr, request) + lent = set(kit.pages(kv, 0)) + lender = attach_in_place(mgr) + for lend in (lender.lend_read, lender.lend_write)[:leases]: + lend(request, 0, END).release() + mgr.free_resources(request) + assert kit.closed(kv), "freed at once" + assert kit.taken_by_others(mgr, lent), "the freed pages go to the next requests" + kept = len(kit.retained()) + spy = mgr.impl = kit.ShutdownSpy(mgr.impl) + warnings = kit.lender_warnings(monkeypatch) + mgr.shutdown() + assert spy.shut_down and warnings == [] and len(kit.retained()) == kept + + +# -- references and threads ------------------------------------------------------------------- + + +def check_a_dropped_lease_ends_no_loan(kit, real_manager, monkeypatch): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + ended = [] + end_loan = _lender.InPlace._end_loan + + def spy(self, kv_cache): + ended.append(threading.get_ident()) + return end_loan(self, kv_cache) + + monkeypatch.setattr(_lender.InPlace, "_end_loan", spy) + with real_manager(max_tokens=kit.POOL_TOKENS) as mgr: + request = kit.published(mgr, SOURCE, PROMPT) + kv = kit.kv(mgr, request) + lender = attach_in_place(mgr) + cycle = [lender.lend_read(request, 0, END)] + cycle.append(cycle) # only a collector frees it + mgr.free_resources(request) + del cycle + in_turn(kit, gc.collect, "collector") + assert ended == [], "collecting a lease ended its loan" + assert not kit.closed(kv) and kit.kv(mgr, request) is None + mgr.shutdown() + assert any(o is kv for o in kit.retained()), "the shutdown keeps what is still lent" + + +def test_a_dropped_lease_ends_no_loan_whatever_thread_collects_it(kit, real_manager, monkeypatch): + check_a_dropped_lease_ends_no_loan(kit, real_manager, monkeypatch) + + +def test_the_check_catches_a_lease_whose_collection_ends_its_loan(kit, real_manager, monkeypatch): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + init = _lender._InPlaceLease.__init__ + + def hold_lender_strongly(self, lender, *args, **kwargs): + init(self, lender, *args, **kwargs) + # The collector clears a weak reference inside the garbage before any finalizer runs. + self._lender = lambda: lender + + def release_when_collected(self): + self.release() + + monkeypatch.setattr(_lender._InPlaceLease, "__init__", hold_lender_strongly) + monkeypatch.setattr(_lender._InPlaceLease, "__del__", release_when_collected, raising=False) + with pytest.raises(AssertionError, match="collecting a lease ended its loan"): + check_a_dropped_lease_ends_no_loan(kit, real_manager, monkeypatch) + + +def test_open_leases_keep_neither_the_lender_nor_the_manager_alive(kit): + torch.cuda.init() + gc.collect() + mgr = kit.make_manager() + request = kit.published(mgr, SOURCE, PROMPT) + lender = attach_in_place(mgr) + read = lender.lend_read(request, 0, END) + write = lender.lend_write(request, 0, END) + view = write.poll() + watched = (weakref.ref(mgr), weakref.ref(lender)) + mgr._stream.synchronize() + del mgr, lender + gc.collect() + assert [ref() for ref in watched] == [None, None], "a lease kept its lender or manager alive" + write.mark_arrived(view.row_masks(True)) + for lease in (write, write, read): + lease.release() + gc.collect() + torch.cuda.empty_cache() + + +def test_threads_take_turns_and_a_release_closes_a_freed_cache_on_its_own(kit, real_manager): + with real_manager(max_tokens=kit.POOL_TOKENS) as mgr: + request = kit.published(mgr, SOURCE, PROMPT) + kv = kit.kv(mgr, request) + before = set(threading.enumerate()) + lender = in_turn(kit, lambda: attach_in_place(mgr), "builder")["value"] + + def executor_loop(): + lease = lender.lend_read(request, 0, END) + rows = lease.poll().num_rows + mgr.free_resources(request) + lease.release() # closes the freed cache inside this call, on this thread + return rows, kit.closed(kv) + + looped = in_turn(kit, executor_loop, "executor-loop") + assert looped.get("value") == (ceil_blocks(END), True) + assert set(threading.enumerate()) <= before, "the lender started a thread" + assert "error" not in in_turn(kit, mgr.shutdown, "shutdown") + + +# -- the manager's page-index buffer under a cache that outlives the manager -------------------- +# A cache writes -1 for every block into its manager's host page-index buffer as it closes. A new +# tensor reclaims a freed buffer's address with a canary the cache must neither read nor write. + +FREED = "touched the page-index buffer freed with its manager" + + +def fresh_device(): + torch.cuda.init() + gc.collect() + torch.cuda.empty_cache() + + +def check_a_lease_outliving_its_manager(kit, free_first): + """A lease outlives its manager and lender, collected without a shutdown; ``free_first`` frees + the request first, which detaches its cache from the buffer.""" + fresh_device() + mgr = kit.make_manager() + request = kit.published(mgr, SOURCE, PROMPT) + kv = kit.kv(mgr, request) + blocks = int(kv.num_blocks) + lender = attach_in_place(mgr) + lease = lender.lend_write(request, 0, END) + assert lease.poll() is not None + row = kit.index_row(mgr, request) + own = list(kv.get_base_page_indices(0)[:blocks]) + assert row.values[:blocks] == own and -1 not in own, "the cache writes elsewhere" + if free_first: + mgr.free_resources(request) + watched = [weakref.ref(o) for o in (mgr.host_kv_cache_block_offsets, mgr, lender)] + mgr._stream.synchronize() + del mgr, lender + gc.collect() + assert watched[1]() is None and watched[2]() is None, "the manager or lender was not collected" + canary = None if watched[0]() is not None else kit.reclaim(row) + if watched[0]() is None: + assert canary is not None, "inconclusive: no allocation reused the freed buffer's address" + read = list(kv.get_base_page_indices(0)[:blocks]) + del kv # the lease now holds the cache's last reference + lease.release() + gc.collect() + written = kit.canary_written(canary, row) + assert not (canary is not None and read == [kit.CANARY] * blocks) and not written, ( + f"the lent cache {FREED}: it read {read} and its close wrote {written} (cell, value)" + ) + + +def test_a_lease_outliving_its_manager_touches_no_freed_index_buffer(kit): + check_a_lease_outliving_its_manager(kit, free_first=False) + + +def test_a_freed_request_s_cache_is_detached_before_its_manager_goes(kit): + check_a_lease_outliving_its_manager(kit, free_first=True) + + +def check_a_cache_kept_at_shutdown(kit): + """The manager shuts down with a loan open, which keeps the cache until exit; the lease is + released and the manager collected. Dropping the kept cache, as the process exit does, must + write nothing into the freed buffer.""" + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + fresh_device() + mgr = kit.make_manager() + request = kit.published(mgr, SOURCE, PROMPT) + kv = kit.kv(mgr, request) + blocks = int(kv.num_blocks) + lender = attach_in_place(mgr) + lease = lender.lend_read(request, 0, END) + assert lease.poll() is not None + row = kit.index_row(mgr, request) + impl = mgr.impl + mgr._stream.synchronize() + mgr.shutdown() + assert any(o is kv for o in kit.retained()), "the shutdown keeps the lent cache" + lease.release() + buffer = weakref.ref(mgr.host_kv_cache_block_offsets) + del mgr, lender + gc.collect() + canary = None if buffer() is not None else kit.reclaim(row) + if buffer() is None: + assert canary is not None, "inconclusive: no allocation reused the freed buffer's address" + read = list(kv.get_base_page_indices(0)[:blocks]) + _lender._let_go(kv) # what the process exit does to the keep list + _lender._let_go(impl) + del kv, impl + gc.collect() + written = kit.canary_written(canary, row) + assert not (canary is not None and read == [kit.CANARY] * blocks) and not written, ( + f"the kept cache {FREED}: it read {read} and its close wrote {written} (cell, value)" + ) + + +def test_a_cache_kept_at_shutdown_touches_no_freed_index_buffer(kit): + check_a_cache_kept_at_shutdown(kit) + + +class _Reclaimer: + """A manager attribute set after the buffer and before the lender, so the manager's collection + drops it between the two: it reclaims the freed buffer's address before the lender goes.""" + + def __init__(self, kit, row, out): + self._kit, self._row, self._out = kit, row, out + + def __del__(self): + self._out.append(self._kit.reclaim(self._row)) + + +def check_a_dropped_lease_on_a_collected_manager(kit): + """A lease dropped unreleased leaves its loan with the lender; the manager is collected without + a shutdown, its attributes in order. The cache the lender held closes after the buffer went.""" + fresh_device() + mgr = kit.make_manager() + request = kit.published(mgr, SOURCE, PROMPT) + row = kit.index_row(mgr, request) + out = [] + mgr._reclaimer = _Reclaimer(kit, row, out) + lender = attach_in_place(mgr) + lender.lend_read(request, 0, END) # dropped unreleased + del lender + mgr._stream.synchronize() + watched, buffer = weakref.ref(mgr), weakref.ref(mgr.host_kv_cache_block_offsets) + del mgr + gc.collect() + assert watched() is None + canary = out[0] if out else None + if buffer() is None: + assert canary is not None, "inconclusive: the freed address was not reclaimed" + written = kit.canary_written(canary, row) + assert not written, f"the cache the lender held {FREED}: its close wrote {written}" + + +def test_a_dropped_lease_on_a_collected_manager_touches_no_freed_index_buffer(kit): + check_a_dropped_lease_on_a_collected_manager(kit) + + +def check_a_last_release_after_its_manager(kit, cache_held): + """The manager is collected without a shutdown or a free while a loan is open; the caller still + holds the lender and releases the last lease. The cache closes as the release drops its last + reference or, with ``cache_held``, as the test drops its own.""" + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + fresh_device() + mgr = kit.make_manager() + request = kit.published(mgr, SOURCE, PROMPT) + kv = kit.kv(mgr, request) + blocks = int(kv.num_blocks) + lender = attach_in_place(mgr) + lease = lender.lend_write(request, 0, END) + assert lease.poll() is not None + row = kit.index_row(mgr, request) + own = list(kv.get_base_page_indices(0)[:blocks]) + assert row.values[:blocks] == own and -1 not in own, "the cache writes elsewhere" + buffer, manager = weakref.ref(mgr.host_kv_cache_block_offsets), weakref.ref(mgr) + mgr._stream.synchronize() + del mgr + gc.collect() + assert manager() is None, "the manager was not collected" + out = [] + + def reclaim_if_freed(): + gc.collect() + if buffer() is None and not out: + out.append(kit.reclaim(row)) + + reclaim_if_freed() # a buffer the manager's collection freed + if cache_held: + lease.release() + reclaim_if_freed() + del kv + else: + end_loan = _lender.InPlace._end_loan + + def end_loan_then_reclaim(self, kv_cache): + # Reclaims a buffer freed here, before the release drops the cache's last reference. + end_loan(self, kv_cache) + reclaim_if_freed() + + with pytest.MonkeyPatch.context() as patch: + patch.setattr(_lender.InPlace, "_end_loan", end_loan_then_reclaim) + del kv # the lease and the loan now hold the cache's last references + lease.release() + gc.collect() + canary = out[0] if out else None + if buffer() is None: + assert canary is not None, "inconclusive: no allocation reused the freed buffer's address" + written = kit.canary_written(canary, row) + assert not written, f"the lent cache {FREED}: its close wrote {written} (cell, value)" + + +@pytest.mark.parametrize("cache_held", [False, True], ids=["release_drops_cache", "cache_held"]) +def test_a_last_release_after_its_manager_touches_no_freed_index_buffer(kit, cache_held): + check_a_last_release_after_its_manager(kit, cache_held) + + +@pytest.mark.parametrize("cache_held", [False, True], ids=["release_drops_cache", "cache_held"]) +def test_the_check_catches_a_lender_letting_the_buffer_go_after_its_manager( + kit, monkeypatch, cache_held +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + def let_go_regardless(self): + _lender._let_go(self._index_buffer) + self._index_buffer = None + + monkeypatch.setattr(_lender.InPlace, "_let_go_index_buffer", let_go_regardless) + with pytest.raises(AssertionError, match=FREED): + check_a_last_release_after_its_manager(kit, cache_held) + + +@pytest.mark.parametrize( + "check", + [ + lambda kit: check_a_lease_outliving_its_manager(kit, free_first=False), + check_a_cache_kept_at_shutdown, + check_a_dropped_lease_on_a_collected_manager, + lambda kit: check_a_last_release_after_its_manager(kit, cache_held=False), + lambda kit: check_a_last_release_after_its_manager(kit, cache_held=True), + ], + ids=[ + "lease_outlives", + "kept_at_shutdown", + "dropped_lease", + "last_release_drops_cache", + "last_release_cache_held", + ], +) +def test_the_checks_catch_a_lender_not_keeping_the_index_buffer(kit, monkeypatch, check): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + monkeypatch.setattr(_lender.InPlace, "_keep_index_buffer", lambda self, manager: None) + with pytest.raises(AssertionError, match=FREED): + check(kit) + + +def check_the_index_buffer_is_kept_only_while_lent(kit, real_manager): + with real_manager(max_tokens=kit.POOL_TOKENS) as mgr: + request = kit.published(mgr, SOURCE, PROMPT) + before = len(kit.retained()) + buffer = mgr.host_kv_cache_block_offsets + lender = attach_in_place(mgr) + assert not any(o is buffer for o in kit.retained()), "kept with no loan open" + leases = [lender.lend_read(request, 0, END), lender.lend_write(request, 0, END)] + assert any(o is buffer for o in kit.retained()), "not kept while a loan is open" + leases[0].release() + assert any(o is buffer for o in kit.retained()), "let go with a loan still open" + leases[1].release() + assert len(kit.retained()) == before, "kept after the last loan ended" + + +def test_the_index_buffer_is_kept_only_while_a_loan_is_open(kit, real_manager): + check_the_index_buffer_is_kept_only_while_lent(kit, real_manager) + + +def test_the_check_catches_a_lender_keeping_the_index_buffer_for_good( + kit, real_manager, monkeypatch +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + monkeypatch.setattr(_lender.InPlace, "_let_go_index_buffer", lambda self: None) + with pytest.raises(AssertionError, match="kept after the last loan ended"): + check_the_index_buffer_is_kept_only_while_lent(kit, real_manager) diff --git a/tests/unittest/_torch/executor/kv_cache/sharing/test_layout_identity.py b/tests/unittest/_torch/executor/kv_cache/sharing/test_layout_identity.py new file mode 100644 index 000000000000..f66d6e949e5c --- /dev/null +++ b/tests/unittest/_torch/executor/kv_cache/sharing/test_layout_identity.py @@ -0,0 +1,1261 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""``_layout``: ``layout_id`` and part names over hand-built layouts (CPU), ``derive_layout`` over +a stand-in manager scripting TP ranks, pipeline stages and recurrent groups (CPU), and over real +``KVCacheManagerV2`` instances (GPU): what each description comes from and what moves the id.""" + +import dataclasses +import enum +import gc +import hashlib +import json +import os +from contextlib import contextmanager +from types import SimpleNamespace +from typing import Dict, List, Optional, Sequence + +import numpy as np +import pytest +import torch +from utils.util import skip_pre_blackwell, skip_pre_hopper + +from tensorrt_llm._torch.disaggregation.resource.page import MapperKind, RoleLayout +from tensorrt_llm._torch.pyexecutor.kv_cache.sharing._identity import Identity +from tensorrt_llm._torch.pyexecutor.kv_cache.sharing._layout import ( + LAYOUT_ID_VERSION, + BufferDesc, + BufferGeometry, + BufferMapper, + LayerGroupDesc, + Layout, + ShardDesc, + canonical_bytes, + canonical_order, + derive_layout, + global_layer_ids, + layout_document, + layout_id, +) + +gpu = pytest.mark.skipif(not torch.cuda.is_available(), reason="allocates KV cache pools") + +# -- hand-built layouts -------------------------------------------------------------------------- + +# Device pool group 0 has two pools (a 100-byte and a 40-byte slot); pool group 1 is a single-pool +# group for a second layer group. + + +def make_layout( + *, + tokens_per_block=32, + windows=(None, 64), + layers=((0, 2), (1,)), + shard=ShardDesc(), + group_ids=(0, 1), + slot_bytes=(100, 40), +) -> Layout: + g0, g1 = group_ids + pool_groups = {g0: tuple(slot_bytes), g1: (64,)} + layer_groups = ( + LayerGroupDesc( + kind="attention", + window=windows[0], + sink_blocks=0, + pool_group=g0, + layers=layers[0], + shard=shard, + ), + LayerGroupDesc( + kind="attention", + window=windows[1], + sink_blocks=0, + pool_group=g1, + layers=layers[1], + shard=shard, + ), + ) + buffers = [] + for layer in layers[0]: + buffers += [ + BufferDesc(layer, "key", (g0, 0), 50 * (layer // 2), 50), + BufferDesc(layer, "scale", (g0, 1), 20 * (layer // 2), 20), + ] + for layer in layers[1]: + buffers.append(BufferDesc(layer, "key", (g1, 0), 0, 64)) + return Layout(tokens_per_block, pool_groups, layer_groups, tuple(buffers)) + + +def _replace_buffers(layout, **change): + return dataclasses.replace( + layout, + buffers=tuple( + dataclasses.replace(b, **change) if b.role == "scale" else b for b in layout.buffers + ), + ) + + +@pytest.mark.cpu_only +def test_layout_id_is_a_stable_sha256(): + first, second = layout_id(make_layout()), layout_id(make_layout()) + assert len(first) == 32 + assert first == second + + +@pytest.mark.cpu_only +def test_layout_id_ignores_local_pool_group_numbering(): + assert layout_id(make_layout(group_ids=(0, 1))) == layout_id(make_layout(group_ids=(7, 3))) + + +@pytest.mark.cpu_only +def test_layout_id_ignores_the_order_of_local_layer_groups(): + layout = make_layout() + swapped = dataclasses.replace(layout, layer_groups=layout.layer_groups[::-1]) + assert layout_id(swapped) == layout_id(layout) + + +@pytest.mark.cpu_only +def test_layout_id_ignores_sinks(): + layout = make_layout() + lg0 = dataclasses.replace(layout.layer_groups[0], sink_blocks=3) + changed = dataclasses.replace(layout, layer_groups=(lg0,) + layout.layer_groups[1:]) + assert layout_id(changed) == layout_id(layout) + + +@pytest.mark.cpu_only +@pytest.mark.parametrize( + "change", + [dict(windows=(None, 128)), dict(windows=(64, 64)), dict(shard=ShardDesc(2, 1))], + ids=["window", "full_to_window", "shard"], +) +def test_layout_id_leaves_model_semantics_and_the_shard_position_out(change): + """Windows say which blocks exist, not how their bytes read; the shard position rides in the + name, so every rank computes the same ``layout_id``.""" + assert layout_id(make_layout(**change)) == layout_id(make_layout()) + + +@pytest.mark.cpu_only +@pytest.mark.parametrize( + "change", + [ + lambda: make_layout(tokens_per_block=64), + lambda: make_layout(layers=((0, 1), (2,))), + lambda: make_layout(slot_bytes=(104, 40)), + lambda: _replace_buffers(make_layout(), mapper=BufferMapper.REPLICATED), + lambda: _replace_buffers(make_layout(), geometry=BufferGeometry(bytes_per_head=10)), + lambda: _replace_buffers(make_layout(), geometry=BufferGeometry(section_bytes=(15, 5))), + lambda: _replace_buffers(make_layout(), transfer=False), + lambda: _replace_buffers(make_layout(), expansion=2), + lambda: _replace_buffers(make_layout(), role="scale2"), + lambda: dataclasses.replace( + make_layout(), + layer_groups=( + dataclasses.replace(make_layout().layer_groups[0], kind="state"), + make_layout().layer_groups[1], + ), + ), + ], + ids=[ + "tokens_per_block", + "layer_membership", + "slot_width", + "mapper", + "head_geometry", + "sections", + "transfer", + "expansion", + "role", + "kind", + ], +) +def test_layout_id_changes_with_every_input_that_decides_how_bytes_read(change): + assert layout_id(change()) != layout_id(make_layout()) + + +@pytest.mark.cpu_only +def test_layout_id_changes_with_a_buffer_s_position_in_the_transfer_bytes(): + layout = make_layout() + moved = tuple( + dataclasses.replace(b, offset=b.offset + 10) if b.role == "scale" else b + for b in layout.buffers + ) + assert layout_id(dataclasses.replace(layout, buffers=moved)) != layout_id(layout) + + +@pytest.mark.cpu_only +def test_canonical_order_sorts_groups_by_smallest_layer(): + layout = make_layout(layers=((3, 5), (1, 4))) + assert canonical_order([d.layers for d in layout.layer_groups]) == [1, 0] + assert canonical_order([d.layers for d in layout.layer_groups[::-1]]) == [0, 1] + # Ties and empty groups keep the local order; empty groups go last. + assert canonical_order([(), (2,), (2,), (0, 9)]) == [3, 1, 2, 0] + + +@pytest.mark.cpu_only +def test_canonical_order_is_a_bijection(): + order = canonical_order([(9,), (2, 11), (5,)]) + assert sorted(order) == [0, 1, 2] + + +def _identity(layout, scope=b"model"): + return Identity( + scope, + layout_id(layout), + [group.layers for group in layout.layer_groups], + [(group.shard.count, group.shard.index) for group in layout.layer_groups], + ) + + +def _part_names(layout, scope=b"model"): + identity = _identity(layout, scope) + return [identity.part_name([lg]) for lg in range(len(layout.layer_groups))] + + +@pytest.mark.cpu_only +def test_part_names_follow_canonical_layer_groups_not_local_order(): + layout = make_layout(layers=((3, 5), (1, 4))) + swapped = dataclasses.replace(layout, layer_groups=layout.layer_groups[::-1]) + prefix = layout_id(layout).hex()[:16] + assert _part_names(layout) == [f"{prefix}:lg1", f"{prefix}:lg0"] + assert _part_names(swapped) == [f"{prefix}:lg0", f"{prefix}:lg1"] + assert _identity(layout).part_name([1, 0]) == f"{prefix}:lg0+1" + assert _part_names(layout, b"other weights") == _part_names(layout) + + +@pytest.mark.cpu_only +def test_part_names_agree_between_layouts_built_apart_and_differ_otherwise(): + assert _part_names(make_layout()) == _part_names(make_layout()) + # Same slot sizes, another block layout: a peer's size check would pass, the names must not. + other = make_layout(tokens_per_block=16) + assert not set(_part_names(other)) & set(_part_names(make_layout())) + + +@pytest.mark.cpu_only +def test_layout_id_is_the_hash_of_the_documented_version_2_encoding(): + """The canonical encoding, spelled out here byte for byte so a change to it cannot pass + unnoticed.""" + layout = make_layout(layers=((0,), (1,))) + expected = ( + '{"buffers":[' + '{"layer":0,"offset":0,"role":"key","size":50},' + '{"layer":0,"offset":100,"role":"scale","size":20},' + '{"layer":1,"offset":0,"role":"key","size":64}],' + '"groups":[{"layers":[0],"slot_bytes":[100,40]},{"layers":[1],"slot_bytes":[64]}],' + '"tokens_per_block":32,"v":2}' + ).encode() + assert LAYOUT_ID_VERSION == 2 + assert canonical_bytes(layout_document(layout)) == expected + assert layout_id(layout) == hashlib.sha256(expected).digest() + + +@pytest.mark.cpu_only +def test_named_keys_at_their_defaults_are_left_out_and_others_are_written(): + layout = make_layout(layers=((0,), (1,))) + plain = layout_document(layout)["buffers"][1] + assert set(plain) == {"layer", "role", "offset", "size"} + changed = layout_document( + _replace_buffers( + layout, + mapper=BufferMapper.SECTIONED, + geometry=BufferGeometry(section_bytes=(15, 5)), + transfer=False, + expansion=4, + ) + )["buffers"][1] + assert changed["mapper"] == "SECTIONED" + assert changed["geometry"] == {"section_bytes": [15, 5]} + assert changed["transfer"] is False + assert changed["expansion"] == 4 + state = dataclasses.replace( + layout, + layer_groups=(dataclasses.replace(layout.layer_groups[0], kind="state"),) + + layout.layer_groups[1:], + ) + assert layout_document(state)["groups"][0]["kind"] == "state" + assert "kind" not in layout_document(layout)["groups"][0] + + +@pytest.mark.cpu_only +def test_enums_encode_by_name_and_the_hnd_alias_is_indexed(): + hnd = _replace_buffers(make_layout(), mapper=MapperKind.HND) + indexed = _replace_buffers(make_layout(), mapper=BufferMapper.INDEXED) + assert layout_id(hnd) == layout_id(indexed) == layout_id(make_layout()) + + class Color(enum.IntEnum): + RED = 3 + + assert json.loads(canonical_bytes({"c": Color.RED, "n": None, "l": [None, 1]})) == { + "c": "RED", + "l": [None, 1], + } + + +@pytest.mark.cpu_only +def test_floats_encode_as_their_binary64_bits(): + assert canonical_bytes({"x": 1.0}) == b'{"x":"f64:3ff0000000000000"}' + assert canonical_bytes({"x": -0.0}) != canonical_bytes({"x": 0.0}) + + +@pytest.mark.cpu_only +def test_buffers_are_described_by_global_layer_so_the_listing_order_does_not_matter(): + layout = make_layout() + shuffled = dataclasses.replace(layout, buffers=layout.buffers[::-1]) + assert layout_id(shuffled) == layout_id(layout) + + +@pytest.mark.cpu_only +def test_a_document_refuses_what_the_layout_does_not_hold(): + layout = make_layout() + stray = dataclasses.replace(layout.layer_groups[0], pool_group=9) + with pytest.raises(ValueError, match="maps to pool group 9"): + layout_document( + dataclasses.replace(layout, layer_groups=(stray,) + layout.layer_groups[1:]) + ) + outside = layout.buffers + (BufferDesc(7, "key", (9, 0), 0, 8),) + with pytest.raises(ValueError, match="not in the layout"): + layout_document(dataclasses.replace(layout, buffers=outside)) + twice = layout.buffers + (layout.buffers[0],) + with pytest.raises(ValueError, match="appears twice"): + layout_document(dataclasses.replace(layout, buffers=twice)) + + +# -- geometry and buffers ------------------------------------------------------------------------ + + +@pytest.mark.cpu_only +def test_a_geometry_that_does_not_fit_its_buffer_is_refused(): + with pytest.raises(ValueError, match="sum"): + BufferDesc(0, "conv", (0, 0), 0, 30, geometry=BufferGeometry(section_bytes=(10, 10))) + with pytest.raises(ValueError, match="heads"): + BufferDesc(0, "k", (0, 0), 0, 30, geometry=BufferGeometry(bytes_per_head=8)) + with pytest.raises(ValueError, match="heads"): + BufferDesc(0, "k", (0, 0), 0, 32, geometry=BufferGeometry(bytes_per_head=8, num_heads=2)) + + +@pytest.mark.cpu_only +def test_an_expansion_must_divide_the_block(): + with pytest.raises(ValueError, match="positive"): + BufferDesc(0, "k", (0, 0), 0, 32, expansion=0) + with pytest.raises(ValueError, match="does not divide"): + _replace_buffers(make_layout(), expansion=3) + + +@pytest.mark.cpu_only +@pytest.mark.parametrize("shard", [(2, 2), (0, 0), (2, -1)]) +def test_a_shard_is_a_share_of_its_count(shard): + with pytest.raises(ValueError, match="shard"): + ShardDesc(*shard) + + +@pytest.mark.cpu_only +@pytest.mark.parametrize( + "sections", + [(np.int64(4), np.int64(4)), np.array([4, 4]), [4, 4], (4, 4)], + ids=["numpy-ints", "numpy-array", "list", "tuple"], +) +def test_a_validated_geometry_holds_its_sections_as_a_tuple_of_ints(sections): + checked = BufferGeometry(section_bytes=sections, bytes_per_head=2).validate() + assert checked == BufferGeometry(section_bytes=(4, 4), bytes_per_head=2) + assert type(checked.section_bytes) is tuple + assert all(type(s) is int for s in checked.section_bytes) + + +@pytest.mark.cpu_only +@pytest.mark.parametrize( + "fields", [{"section_bytes": (4, 0)}, {"bytes_per_head": 0}, {"num_heads": -1}] +) +def test_validation_refuses_sizes_that_are_not_positive(fields): + with pytest.raises(ValueError, match="positive"): + BufferGeometry(**fields).validate() + + +# -- a stand-in manager -------------------------------------------------------------------------- + +TPB = 4 +HEAD_BYTES = TPB * 8 * 2 # 8 fp16 values per token and head + + +class AttentionLayerConfig(SimpleNamespace): + pass + + +class SsmLayerConfig(SimpleNamespace): + """Named like the manager's state layer config: ``derive_layout`` tells groups apart by it.""" + + +def buffer(role, override=None): + return SimpleNamespace(role=role, tokens_per_block_override=override) + + +def coalesced(size, members): + return SimpleNamespace( + single_buffer_size=size, + buffer_ids=[SimpleNamespace(layer_id=layer, role=role) for layer, role in members], + ) + + +class Manager: + """One layer group per entry of ``groups``: ``(layer config class, [internal layer ids], + [(role, size, override)])``; group ``i`` draws from pool group ``pool_group_of[i]`` (its own by + default), one pool per role, and has window ``windows[i]`` and ``sinks[i]`` sink tokens.""" + + def __init__( + self, + groups, + *, + tp_size=1, + tp_rank=0, + attention_dp=False, + num_kv_heads=4, + pp_layers: Optional[Sequence[int]] = None, + mappers: Optional[Dict[str, MapperKind]] = None, + layouts: Optional[Dict[str, RoleLayout]] = None, + ignored=(), + pool_group_of: Optional[Sequence[int]] = None, + windows: Optional[Sequence[Optional[int]]] = None, + sinks: Optional[Sequence[Optional[int]]] = None, + ): + heads_per_rank = -(-num_kv_heads // (1 if attention_dp else tp_size)) + internal = sorted({lid for _, lids, _ in groups for lid in lids}) + pool_group_of = list(pool_group_of if pool_group_of is not None else range(len(groups))) + windows = list(windows or [None] * len(groups)) + sinks = list(sinks or [None] * len(groups)) + self.tokens_per_block = TPB + self.mapping = SimpleNamespace( + tp_size=tp_size, tp_rank=tp_rank, enable_attention_dp=attention_dp + ) + self.num_kv_heads = num_kv_heads + self.num_kv_heads_per_layer = [heads_per_rank] * len(internal) + self.pp_layers = list(pp_layers if pp_layers is not None else internal) + self.mappers = {"all": MapperKind.INDEXED, **(mappers or {})} + self.layouts = layouts or {} + self.ignored = frozenset(ignored) + layers: List[object] = [] + descs = {} + for lg, (cls, lids, roles) in enumerate(groups): + for lid in lids: + layers.append( + cls( + layer_id=lid, + buffers=[buffer(role, override) for role, _, override in roles], + window_size=windows[lg], + num_sink_tokens=sinks[lg], + ) + ) + pools = [coalesced(size, [(lid, role) for lid in lids]) for role, size, _ in roles] + variant = SimpleNamespace(layer_group_id=lg, coalesced_buffers=pools) + g = pool_group_of[lg] + if g in descs: + descs[g].slot_desc.variants.append(variant) + continue + descs[g] = SimpleNamespace( + pool_group_index=g, + num_slots=8, + pools=[ + SimpleNamespace( + pool_index=i, + base_address=0x1000 * (g + 1) + 0x100 * i, + slot_bytes=c.single_buffer_size * len(lids), + ) + for i, c in enumerate(pools) + ], + slot_desc=SimpleNamespace(variants=[variant]), + ) + self.impl = SimpleNamespace( + layer_grouping=[list(lids) for _, lids, _ in groups], + pool_group_descs=list(descs.values()), + get_life_cycle_pool_group_indices=lambda: list(pool_group_of), + init_config=SimpleNamespace(layers=layers), + ) + + def get_disagg_role_mapper_kinds(self): + return self.mappers + + def get_disagg_role_layouts(self): + return self.layouts + + def get_disagg_ignored_roles(self): + return self.ignored + + +def kv_group(lids=(0, 1), heads_per_rank=1, extra=()): + size = HEAD_BYTES * heads_per_rank + return (AttentionLayerConfig, list(lids), [("key", size, None), ("value", size, None), *extra]) + + +def attention(tp_size, tp_rank, num_kv_heads, **kwargs): + heads = -(-num_kv_heads // tp_size) + return Manager( + [kv_group(heads_per_rank=heads)], + tp_size=tp_size, + tp_rank=tp_rank, + num_kv_heads=num_kv_heads, + **kwargs, + ) + + +def recurrent_state(): + return Manager( + [(SsmLayerConfig, [0, 1], [("conv_state", 96, None), ("ssm_state", 128, None)])], + tp_size=2, + tp_rank=1, + mappers={"conv_state": MapperKind.SECTIONED}, + layouts={ + "conv_state": RoleLayout(section_bytes=(32, 32, 32)), + "ssm_state": RoleLayout(bytes_per_head=32), + }, + ) + + +TOKEN_MAJOR_SHARES = "head-major and token-major K/V share a layout_id" + + +def check_token_major_kv_has_its_own_layout_id(): + """A manager declaring token-major K/V (NHD) and one declaring the default head-major K/V (HND) + lay a block's bytes out differently, so they never share a ``layout_id``.""" + head_major = derive_layout(attention(1, 0, 4)).layout + token_major = derive_layout(attention(1, 0, 4, mappers={"all": MapperKind.NHD})).layout + assert layout_id(head_major) != layout_id(token_major), TOKEN_MAJOR_SHARES + assert {b.mapper for b in token_major.buffers} == {BufferMapper.NHD} + assert {b.mapper for b in head_major.buffers} == {BufferMapper.INDEXED} + hnd = derive_layout(attention(1, 0, 4, mappers={"all": MapperKind.HND})).layout + assert layout_id(hnd) == layout_id(head_major), "HND is the default head-major layout" + + +@pytest.mark.cpu_only +def test_token_major_kv_has_its_own_layout_id(): + check_token_major_kv_has_its_own_layout_id() + + +@pytest.mark.cpu_only +def test_the_check_catches_a_derivation_ignoring_the_declared_layout(monkeypatch): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _layout + + declared = _layout._declared + + def head_major_always(manager, getter): + return {} if getter == "get_disagg_role_mapper_kinds" else declared(manager, getter) + + monkeypatch.setattr(_layout, "_declared", head_major_always) + with pytest.raises(AssertionError, match=TOKEN_MAJOR_SHARES): + check_token_major_kv_has_its_own_layout_id() + + +@pytest.mark.cpu_only +def test_the_check_catches_a_layout_id_leaving_the_arrangement_out(monkeypatch): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _layout + + document = _layout._buffer_document + + def without_mapper(layout, buffer): + return {k: v for k, v in document(layout, buffer).items() if k != "mapper"} + + monkeypatch.setattr(_layout, "_buffer_document", without_mapper) + with pytest.raises(AssertionError, match=TOKEN_MAJOR_SHARES): + check_token_major_kv_has_its_own_layout_id() + + +@pytest.mark.cpu_only +def test_a_single_kv_head_is_the_same_content_on_every_rank(): + """DeepSeek-V4 at TP8 without attention DP: one KV head, repeated on every rank. Every rank + describes the group as whole, so all of them compute one layout and one namespace.""" + geometries = [derive_layout(attention(8, r, num_kv_heads=1)) for r in range(8)] + assert all(geo.layout.layer_groups[0].shard == ShardDesc() for geo in geometries) + assert all(geo.shards == ((1, 0),) for geo in geometries) + assert len({layout_id(geo.layout) for geo in geometries}) == 1 + + +@pytest.mark.cpu_only +def test_head_sharded_kv_names_the_share_each_rank_holds(): + geometries = [derive_layout(attention(4, r, num_kv_heads=8)) for r in range(4)] + assert [geo.layout.layer_groups[0].shard for geo in geometries] == [ + ShardDesc(4, r) for r in range(4) + ] + assert [geo.shards for geo in geometries] == [((4, r),) for r in range(4)] + # The byte layout is the same on every rank; the share is part of the object name. + assert len({layout_id(geo.layout) for geo in geometries}) == 1 + + +@pytest.mark.cpu_only +def test_repeated_heads_are_shared_by_the_ranks_that_repeat_them(): + """Two KV heads over four ranks: ranks 0 and 1 hold head 0, ranks 2 and 3 head 1, so there are + two shares, not four.""" + shards = [ + derive_layout(attention(4, r, num_kv_heads=2)).layout.layer_groups[0].shard + for r in range(4) + ] + assert shards == [ShardDesc(2, 0), ShardDesc(2, 0), ShardDesc(2, 1), ShardDesc(2, 1)] + + +@pytest.mark.cpu_only +def test_attention_dp_holds_whole_heads(): + layout = derive_layout(attention(4, 3, num_kv_heads=8, attention_dp=True)).layout + assert layout.layer_groups[0].shard == ShardDesc() + + +@pytest.mark.cpu_only +def test_a_group_of_replicated_buffers_is_whole_and_one_head_sharded_buffer_makes_it_a_share(): + index = ("index_key", 64, None) + replicated_only = Manager( + [(AttentionLayerConfig, [0], [index])], + tp_size=2, + tp_rank=1, + num_kv_heads=8, + mappers={"index_key": MapperKind.REPLICATED}, + ) + assert derive_layout(replicated_only).layout.layer_groups[0].shard == ShardDesc() + mixed = Manager( + [kv_group(heads_per_rank=4, extra=[index])], + tp_size=2, + tp_rank=1, + num_kv_heads=8, + mappers={"index_key": MapperKind.REPLICATED}, + ) + assert derive_layout(mixed).layout.layer_groups[0].shard == ShardDesc(2, 1) + + +@pytest.mark.cpu_only +def test_local_only_buffers_do_not_decide_the_identity(): + scratch = ("scratch", 32, None) + manager = Manager( + [(AttentionLayerConfig, [0], [("index_key", 64, None), scratch])], + tp_size=2, + tp_rank=1, + mappers={"index_key": MapperKind.REPLICATED}, + ignored=("scratch",), + ) + layout = derive_layout(manager).layout + assert layout.layer_groups[0].shard == ShardDesc() + assert {b.role: b.transfer for b in layout.buffers} == {"index_key": True, "scratch": False} + + +@pytest.mark.cpu_only +def test_groups_of_one_manager_are_judged_one_by_one(): + manager = Manager( + [ + kv_group(lids=(0,), heads_per_rank=4), + (AttentionLayerConfig, [1], [("index_key", 64, None)]), + ], + tp_size=2, + tp_rank=0, + num_kv_heads=8, + mappers={"index_key": MapperKind.REPLICATED}, + ) + groups = derive_layout(manager).layout.layer_groups + assert [g.shard for g in groups] == [ShardDesc(2, 0), ShardDesc()] + + +@pytest.mark.cpu_only +def test_recurrent_state_is_described_and_sharded_by_rank(): + geometry = derive_layout(recurrent_state()) + layout = geometry.layout + (group,) = layout.layer_groups + assert group.kind == "state" and group.shard == ShardDesc(2, 1) + assert geometry.recurrent == (True,) and geometry.windows == (None,) + by_role = {b.role: b for b in layout.buffers} + assert by_role["conv_state"].mapper is BufferMapper.SECTIONED + assert by_role["conv_state"].geometry == BufferGeometry(section_bytes=(32, 32, 32)) + assert by_role["ssm_state"].geometry == BufferGeometry(bytes_per_head=32, num_heads=4) + + +@pytest.mark.cpu_only +def test_head_geometry_is_described_per_buffer(): + layout = derive_layout(attention(2, 0, num_kv_heads=8)).layout + for b in layout.buffers: + assert b.mapper is BufferMapper.INDEXED + assert b.geometry == BufferGeometry(bytes_per_head=HEAD_BYTES, num_heads=4) + + +@pytest.mark.cpu_only +def test_layers_are_global_ids_stable_across_pipeline_stages(): + """Stage 2 of a pipeline holds model layers 10 and 11 as internal layers 0 and 1.""" + manager = Manager([kv_group()], pp_layers=[10, 11]) + geometry = derive_layout(manager) + assert geometry.layout.layer_groups[0].layers == (10, 11) + assert geometry.layers == ((10, 11),) + assert {b.layer for b in geometry.layout.buffers} == {10, 11} + assert global_layer_ids(manager, [1, 0]) == [11, 10] + + +@pytest.mark.cpu_only +def test_virtual_layers_are_numbered_by_model_layer_and_attention_type(): + """Two attention types over three enum members: model layer ``m`` of type ``t`` is + ``3 * m + t``, whichever types this stage holds.""" + + class AttentionType(enum.Enum): + SLIDING = 0 + COMPRESSED = 1 + INDEXER = 2 + + manager = Manager([kv_group(lids=(0, 1, 2))]) + manager._layer_attn_to_layer_id = { + (4, AttentionType.SLIDING): 0, + (4, AttentionType.COMPRESSED): 1, + (5, AttentionType.SLIDING): 2, + } + assert global_layer_ids(manager, [0, 1, 2]) == [12, 13, 15] + assert derive_layout(manager).layers == ((12, 13, 15),) + + +@pytest.mark.cpu_only +def test_declared_global_layer_ids_name_extra_internal_layers(): + """A manager whose internal layers outnumber its model layers (FP4 MLA keeps its tail on + extra ones) declares the ids per layer group; they are used as declared, in group order.""" + manager = Manager( + [kv_group(lids=(1, 0)), (AttentionLayerConfig, [2, 3], [("tail", 64, None)])], + pp_layers=[0, 1], + mappers={"tail": MapperKind.REPLICATED}, + ) + declared = {0: [2, 0], 1: [1, 3]} + manager.get_disagg_global_layer_ids = lambda lg: declared[lg] + assert global_layer_ids(manager, [0, 1, 2, 3]) == [0, 2, 1, 3] + layout = derive_layout(manager).layout + assert [g.layers for g in layout.layer_groups] == [(0, 2), (1, 3)] + assert {(b.layer, b.role) for b in layout.buffers if b.role == "tail"} == { + (1, "tail"), + (3, "tail"), + } + manager.get_disagg_global_layer_ids = lambda lg: declared[lg][:1] + with pytest.raises(ValueError, match="declared global ids"): + global_layer_ids(manager, [0]) + + +@pytest.mark.cpu_only +def test_a_buffer_with_its_own_tokens_per_block_is_expanded(): + manager = Manager( + [kv_group(extra=[("index_key", 64, 2)])], mappers={"index_key": MapperKind.REPLICATED} + ) + layout = derive_layout(manager).layout + assert {b.role: b.expansion for b in layout.buffers} == {"key": 1, "value": 1, "index_key": 2} + plain = derive_layout( + Manager( + [kv_group(extra=[("index_key", 64, None)])], + mappers={"index_key": MapperKind.REPLICATED}, + ) + ).layout + assert layout_id(layout) != layout_id(plain) + + +@pytest.mark.cpu_only +def test_the_manager_layout_holds_per_group_and_per_pool_group_facts(): + """Two attention groups share pool group 3, each with its own window and sinks; a state group + has pool group 1 and neither window nor sinks.""" + manager = Manager( + [ + kv_group(lids=(0,)), + kv_group(lids=(1,)), + (SsmLayerConfig, [2], [("conv_state", 96, None)]), + ], + pool_group_of=[3, 3, 1], + windows=[None, 8, 16], + sinks=[None, 5, 4], + ) + geometry = derive_layout(manager) + assert geometry.num_layer_groups == 3 + assert geometry.tokens_per_block == TPB + assert geometry.pool_group_of == (3, 3, 1) + assert [d.pool_group for d in geometry.layout.layer_groups] == [3, 3, 1] + assert geometry.pool_groups == (1, 3) + assert geometry.windows == (None, 8, None) + assert geometry.sink_blocks == (0, 2, 0) # 5 sink tokens take two 4-token blocks + assert geometry.recurrent == (False, False, True) + assert geometry.layers == ((0,), (1,), (2,)) + assert geometry.shards == ((1, 0),) * 3 + assert dict(geometry.page_bytes) == {3: 2 * HEAD_BYTES, 1: 96} + assert dict(geometry.layout.pool_groups) == {3: (HEAD_BYTES, HEAD_BYTES), 1: (96,)} + pools = { + g: [(p.group, p.index, p.base, p.slot_bytes, p.num_slots) for p in ps] + for g, ps in geometry.device_pools.items() + } + assert pools == { + 3: [(3, 0, 0x4000, HEAD_BYTES, 8), (3, 1, 0x4100, HEAD_BYTES, 8)], + 1: [(1, 0, 0x2000, 96, 8)], + } + + +@pytest.mark.cpu_only +def test_a_manager_whose_parts_disagree_is_refused(): + no_config = Manager([kv_group()]) + no_config.impl.init_config.layers = [] + with pytest.raises(ValueError, match="has no layer config"): + derive_layout(no_config) + miscounted = Manager([kv_group()]) + miscounted.impl.get_life_cycle_pool_group_indices = lambda: [0, 0] + with pytest.raises(ValueError, match="pool-group indices"): + derive_layout(miscounted) + poolless = Manager([kv_group()]) + poolless.impl.get_life_cycle_pool_group_indices = lambda: [5] + with pytest.raises(ValueError, match="which has no pools"): + derive_layout(poolless) + with pytest.raises(ValueError, match="does not divide"): + derive_layout(Manager([kv_group(extra=[("index_key", 64, 3)])])) + + +@pytest.mark.cpu_only +@pytest.mark.parametrize( + "build,expected", + [ + ( + lambda: Manager( + [kv_group(heads_per_rank=4, extra=[("index_key", 64, 2)])], + tp_size=2, + tp_rank=1, + num_kv_heads=8, + mappers={"index_key": MapperKind.REPLICATED}, + ), + "069d747522257b726c6b28d917cb2c260350e26c55550bb687bea7350a76f51c", + ), + ( + recurrent_state, + "3eba9af0efca5fe1e8701398e86b43ada86f07ba249f483be435925ffeaab6fa", + ), + ], + ids=["head_sharded_with_expanded_replicated_role", "recurrent_state"], +) +def test_the_derived_layout_id_is_pinned(build, expected): + """The derivation and the encoding together: a change here moves every name built on it.""" + assert layout_id(derive_layout(build()).layout).hex() == expected + + +# -- real managers ------------------------------------------------------------------------------- + +REAL_TPB = 32 +WINDOW = 64 + + +def native_pool_groups(mgr): + """``{g: [(base, slot bytes)...]}`` and slot counts, read straight off the manager.""" + pools, counts = {}, {} + for pg in mgr.impl.pool_group_descs: + g = int(pg.pool_group_index) + pools[g] = [ + (int(p.base_address), int(p.slot_bytes)) + for p in sorted(pg.pools, key=lambda p: int(p.pool_index)) + ] + counts[g] = int(pg.num_slots) + return pools, counts + + +def check_pool_groups(mgr, geometry): + layout = geometry.layout + pools, counts = native_pool_groups(mgr) + assert set(layout.pool_groups) == set(pools) + assert geometry.pool_groups == tuple(sorted(pools)) + for g, native in pools.items(): + widths = tuple(w for _, w in native) + assert layout.pool_groups[g] == widths + assert geometry.page_bytes[g] == sum(widths) + # Slot i at base + i * the layout's slot width: a pool is its slots back to back. + assert [(p.base, p.slot_bytes * p.num_slots) for p in geometry.device_pools[g]] == [ + (base, width * counts[g]) for base, width in native + ] + + +def check_buffers(mgr, layout): + """Every (layer, role) of the manager appears once, inside its pool's slot.""" + seen = {} + for b in layout.buffers: + assert (b.layer, b.role) not in seen + seen[(b.layer, b.role)] = b + width = layout.pool_groups[b.pool[0]][b.pool[1]] + assert 0 <= b.offset and b.offset + b.size <= width + internal = sorted({int(bid.layer_id) for bid in mgr.impl.all_buffer_ids}) + global_of = dict(zip(internal, global_layer_ids(mgr, internal))) + native = {(global_of[int(bid.layer_id)], str(bid.role)) for bid in mgr.impl.all_buffer_ids} + assert set(seen) == native + + +@gpu +def test_a_full_attention_cache_is_one_layer_group_over_one_pool_group(real_manager): + with real_manager() as mgr: + geometry = derive_layout(mgr) + layout = geometry.layout + + assert layout.tokens_per_block == geometry.tokens_per_block == REAL_TPB + (group,) = layout.layer_groups + assert group.kind == "attention" + assert group.window is None + assert group.layers == (0, 1) + g = int(mgr.impl.get_life_cycle_pool_group_indices()[0]) + assert group.pool_group == g + assert geometry.pool_group_of == (g,) + assert geometry.windows == (None,) + assert geometry.sink_blocks == (0,) + assert geometry.recurrent == (False,) + assert geometry.shards == ((1, 0),) + check_pool_groups(mgr, geometry) + check_buffers(mgr, layout) + + +@gpu +def test_two_windows_are_two_layer_groups_with_their_own_windows(real_manager): + with real_manager(windows=[WINDOW, 256]) as mgr: + geometry = derive_layout(mgr) + layout = geometry.layout + + assert len(layout.layer_groups) == 2 + native = [int(g) for g in mgr.impl.get_life_cycle_pool_group_indices()] + assert list(geometry.pool_group_of) == native + by_layer = {} + for lg, group in enumerate(layout.layer_groups): + assert group.pool_group == native[lg] + assert geometry.windows[lg] == group.window + for layer in group.layers: + by_layer[layer] = group.window + # Layer 0 slides; layer 1's window reaches max_seq_len and normalizes to full attention. + assert mgr.max_seq_len <= 256 + assert by_layer == {0: WINDOW, 1: None} + check_pool_groups(mgr, geometry) + check_buffers(mgr, layout) + + +@gpu +def test_layout_id_ignores_slot_counts_and_addresses(real_manager): + with real_manager(max_tokens=2048) as small: + small_id = layout_id(derive_layout(small).layout) + small_pools = native_pool_groups(small) + with real_manager(max_tokens=8192) as large: + large_geometry = derive_layout(large) + large_pools = native_pool_groups(large) + assert large_pools[1] != small_pools[1], "the two caches must differ in slot counts" + assert layout_id(large_geometry.layout) == small_id + + +@gpu +def test_layout_id_leaves_the_window_out(real_manager): + """One window for every layer: the same single layer group, only its window differs. The + window says which blocks exist, not how a block's bytes read, so the id stays.""" + with real_manager(windows=[WINDOW]) as mgr: + before = derive_layout(mgr).layout + with real_manager(windows=[2 * WINDOW]) as mgr: + after = derive_layout(mgr).layout + assert [g.window for g in before.layer_groups] != [g.window for g in after.layer_groups] + assert layout_id(after) == layout_id(before) + + +@gpu +def test_buffers_carry_their_head_geometry_and_mapper(real_manager): + """K and V are head-sharded (INDEXED): four heads of ``tokens_per_block * head_dim`` fp16 + values on this rank. One rank, so every layer group is whole.""" + with real_manager() as mgr: + layout = derive_layout(mgr).layout + for b in layout.buffers: + assert b.mapper is BufferMapper.INDEXED + assert b.transfer and b.expansion == 1 + assert b.geometry == BufferGeometry(bytes_per_head=REAL_TPB * 64 * 2, num_heads=4) + assert all(group.shard == ShardDesc() for group in layout.layer_groups) + + +@gpu +def test_layout_id_follows_the_tokens_per_block(real_manager): + with real_manager() as mgr: + before = derive_layout(mgr).layout + with real_manager(tokens_per_block=64) as mgr: + after = derive_layout(mgr).layout + assert len(after.layer_groups) == len(before.layer_groups) + assert layout_id(after) != layout_id(before) + + +@contextmanager +def _manager_of_dtype(dtype): + from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import KVCacheManagerV2 + from tensorrt_llm.bindings.internal.batch_manager import CacheType + from tensorrt_llm.llmapi.llm_args import KvCacheConfig + from tensorrt_llm.mapping import Mapping + + gc.collect() + torch.cuda.empty_cache() + mgr = KVCacheManagerV2( + kv_cache_config=KvCacheConfig(max_tokens=2048, enable_block_reuse=True), + kv_cache_type=CacheType.SELF, + num_layers=2, + num_kv_heads=4, + head_dim=64, + tokens_per_block=REAL_TPB, + max_seq_len=256, + max_batch_size=4, + mapping=Mapping(world_size=1, tp_size=1, rank=0), + dtype=dtype, + vocab_size=32000, + execution_stream=torch.cuda.Stream(), + ) + try: + yield mgr + finally: + mgr._stream.synchronize() + mgr.shutdown() + del mgr + gc.collect() + torch.cuda.empty_cache() + + +@gpu +def test_layout_id_leaves_the_element_type_out(): + """FP16 and BF16 caches differ in no size or offset, so they share a ``layout_id``; what the + bytes mean is for the caller's scope to say.""" + from tensorrt_llm.bindings import DataType + + with _manager_of_dtype(DataType.HALF) as mgr: + fp16 = derive_layout(mgr).layout + with _manager_of_dtype(DataType.BF16) as mgr: + bf16 = derive_layout(mgr).layout + assert layout_id(bf16) == layout_id(fp16) + + +@contextmanager +def fp4_mla(backend): + from tensorrt_llm._torch.attention.backends.fp4_mla import ( + FP4_MLA_ATTENTION_BACKEND_ENV, + FP4_MLA_CUTEDSL_FUSED_V_TRANSPOSE_ENV, + ) + from tensorrt_llm._torch.attention.backends.fp4_mla.cache_manager import Fp4MlaKVCacheManagerV2 + from tensorrt_llm._torch.kimi_k3_cache_policy import KIMI_K3_BF16_KV_LAYERS_ENV + from tensorrt_llm.bindings import DataType + from tensorrt_llm.bindings.internal.batch_manager import CacheType + from tensorrt_llm.llmapi.llm_args import KvCacheConfig + from tensorrt_llm.mapping import Mapping + + env = {FP4_MLA_ATTENTION_BACKEND_ENV: backend, FP4_MLA_CUTEDSL_FUSED_V_TRANSPOSE_ENV: "0"} + saved = {name: os.environ.get(name) for name in (*env, KIMI_K3_BF16_KV_LAYERS_ENV)} + os.environ.update(env) + os.environ.pop(KIMI_K3_BF16_KV_LAYERS_ENV, None) + try: + mgr = Fp4MlaKVCacheManagerV2( + KvCacheConfig( + max_tokens=512, dtype="nvfp4", enable_block_reuse=True, host_cache_size=0 + ), + CacheType.SELFKONLY, + num_layers=3, + num_kv_heads=1, + head_dim=576, + tokens_per_block=128, + max_seq_len=512, + max_batch_size=2, + mapping=Mapping(world_size=1, rank=0, tp_size=1), + dtype=DataType.NVFP4, + max_num_tokens=512, + pretrained_config=SimpleNamespace(kv_lora_rank=512), + ) + try: + yield mgr + finally: + mgr._stream.synchronize() + mgr.shutdown() + finally: + for name, value in saved.items(): + if value is None: + os.environ.pop(name, None) + else: + os.environ[name] = value + + +@gpu +@skip_pre_hopper +def test_fp4_mla_buffers_carry_the_declared_global_layer_ids(): + """The tail lives on internal layers past the model's; in the layout its buffers take the ids + the manager declares (``2 * layer + 1``), next to the cache's ``2 * layer``.""" + with fp4_mla("triton") as mgr: + buffers = {(b.layer, b.role) for b in derive_layout(mgr).layout.buffers} + assert {layer for layer, role in buffers if role == "key"} == {0, 2, 4} + assert {layer for layer, role in buffers if role == "mla_hp_tail"} == {1, 3, 5} + + +M3_TPB = 128 +M3_KV_HEADS = 2 +M3_HEAD_DIM = 128 +M3_INDEX_DIM = 128 +M3_SPARSE_LAYER = 3 +M3_KV_DTYPES = {"nvfp4": "NVFP4", "fp8": "FP8"} + + +@contextmanager +def minimax_m3(kv_dtype, max_tokens=1024): + """A four-layer MiniMax-M3 manager: layers 0 to 2 dense, layer 3 sparse with a bf16 index-K. + With NVFP4 KV only the sparse layer is NVFP4 (packed K/V and block scales); dense layers FP8.""" + from tensorrt_llm._torch.attention.backends.sparse.minimax_m3 import MiniMaxM3KVCacheManagerV2 + from tensorrt_llm.bindings import DataType + from tensorrt_llm.bindings.internal.batch_manager import CacheType + from tensorrt_llm.llmapi.llm_args import KvCacheConfig + from tensorrt_llm.mapping import Mapping + + gc.collect() + torch.cuda.empty_cache() + mgr = MiniMaxM3KVCacheManagerV2( + kv_cache_config=KvCacheConfig( + max_tokens=max_tokens, dtype=kv_dtype, enable_block_reuse=True, host_cache_size=0 + ), + kv_cache_type=CacheType.SELF, + num_layers=4, + num_kv_heads=M3_KV_HEADS, + head_dim=M3_HEAD_DIM, + tokens_per_block=M3_TPB, + max_seq_len=512, + max_batch_size=2, + mapping=Mapping(world_size=1, rank=0, tp_size=1), + dtype=getattr(DataType, M3_KV_DTYPES[kv_dtype]), + vocab_size=32000, + max_num_tokens=512, + sparse_layer_ids=[M3_SPARSE_LAYER], + disable_index_value_layer_ids=[M3_SPARSE_LAYER], + sparse_index_dim=M3_INDEX_DIM, + sparse_attention_config=SimpleNamespace(implementation="triton", indexer_kv_dtype="bf16"), + ) + try: + yield mgr + finally: + mgr._stream.synchronize() + mgr.shutdown() + del mgr + gc.collect() + torch.cuda.empty_cache() + + +def check_minimax_m3_nvfp4(mgr, layout): + """Every buffer once, with the size, mapper and heads its storage has: FP8 K/V on dense layers, + half-width K/V plus one scale byte per 16 elements on the sparse layer, replicated index-K.""" + check_buffers(mgr, layout) + fp8 = (M3_TPB * M3_KV_HEADS * M3_HEAD_DIM, BufferMapper.NHD, M3_KV_HEADS) + packed = (fp8[0] // 2, BufferMapper.NHD, M3_KV_HEADS) + scale = (fp8[0] // 16, BufferMapper.NHD, M3_KV_HEADS) + expected = {(layer, role): fp8 for layer in range(4) for role in ("key", "value")} + expected.update({(M3_SPARSE_LAYER, role): packed for role in ("key", "value")}) + expected.update( + {(M3_SPARSE_LAYER, role): scale for role in ("key_block_scale", "value_block_scale")} + ) + expected[(M3_SPARSE_LAYER, "index_key")] = ( + M3_TPB * M3_INDEX_DIM * 2, + BufferMapper.REPLICATED, + None, + ) + described = { + (b.layer, b.role): (b.size, b.mapper, b.geometry.num_heads) for b in layout.buffers + } + assert described == expected + assert all(b.transfer and b.expansion == 1 for b in layout.buffers) + assert sorted(layer for group in layout.layer_groups for layer in group.layers) == [0, 1, 2, 3] + assert all(group.shard == ShardDesc() for group in layout.layer_groups) + + +def minimax_m3_nvfp4_lies(layout): + """Layouts that each misstate one NVFP4 fact: a lost scale, a scale on a dense layer, a scale + as wide as its data and a head-sharded index-K.""" + scale = next(b for b in layout.buffers if b.role == "key_block_scale") + index = next(b for b in layout.buffers if b.role == "index_key") + + def buffers(bufs): + return dataclasses.replace(layout, buffers=tuple(bufs)) + + def swap(old, new): + return buffers(new if b is old else b for b in layout.buffers) + + return { + "scale dropped": buffers(b for b in layout.buffers if b is not scale), + "scale on a dense layer": buffers((*layout.buffers, dataclasses.replace(scale, layer=0))), + "scale as wide as its data": swap( + scale, dataclasses.replace(scale, size=scale.size * 8, geometry=BufferGeometry()) + ), + "index-K head-sharded": swap(index, dataclasses.replace(index, mapper=BufferMapper.NHD)), + } + + +@gpu +@skip_pre_hopper +def test_minimax_m3_nvfp4_describes_its_packed_kv_and_block_scales(): + with minimax_m3("nvfp4") as mgr: + geometry = derive_layout(mgr) + check_pool_groups(mgr, geometry) + check_minimax_m3_nvfp4(mgr, geometry.layout) + for name, lie in minimax_m3_nvfp4_lies(geometry.layout).items(): + with pytest.raises(AssertionError): + check_minimax_m3_nvfp4(mgr, lie) + pytest.fail(f"the {name} lie went unnoticed") + + +def check_minimax_m3_ids(builds, id_of): + """``id_of(layout, slot_counts)`` is one id for every NVFP4 build and another for FP8.""" + ids = {} + for kv_dtype, layout, counts in builds: + ids.setdefault(kv_dtype, set()).add(id_of(layout, counts)) + assert len(ids["nvfp4"]) == 1, "the NVFP4 builds disagree" + assert ids["nvfp4"] != ids["fp8"], "NVFP4 and FP8 share an id" + + +@gpu +@skip_pre_hopper +def test_minimax_m3_nvfp4_and_fp8_caches_have_their_own_layout_ids(): + """NVFP4 builds that differ only in slot count agree; an FP8 cache of the same model reads its + bytes differently.""" + builds = [] + for kv_dtype, max_tokens in (("nvfp4", 1024), ("fp8", 1024), ("nvfp4", 16384)): + with minimax_m3(kv_dtype, max_tokens) as mgr: + builds.append((kv_dtype, derive_layout(mgr).layout, native_pool_groups(mgr)[1])) + assert builds[0][2] != builds[2][2], "the two NVFP4 caches must differ in slot counts" + check_minimax_m3_ids(builds, lambda layout, counts: layout_id(layout)) + lies = { + "one id for every cache": ( + lambda layout, counts: bytes(32), + "NVFP4 and FP8 share an id", + ), + "an id over slot counts": ( + lambda layout, counts: layout_id(layout) + repr(sorted(counts.items())).encode(), + "the NVFP4 builds disagree", + ), + } + for name, (lie, caught_by) in lies.items(): + with pytest.raises(AssertionError, match=caught_by): + check_minimax_m3_ids(builds, lie) + pytest.fail(f"{name} went unnoticed") + + +@gpu +@skip_pre_blackwell +def test_the_deepseek_v4_layout_covers_every_virtual_layer_and_pool(deepseek_v4_manager): + with deepseek_v4_manager() as mgr: + geometry = derive_layout(mgr) + layout = geometry.layout + + life_cycles = mgr._life_cycle_by_layer_group() + assert len(layout.layer_groups) == len(life_cycles) >= 2 + assert list(geometry.windows) == [lc.window_size for lc in life_cycles] + assert None in geometry.windows and any(w is not None for w in geometry.windows) + + # Several pools per group: a page is their slots back to back. + native = {int(pg.pool_group_index): pg for pg in mgr.impl.pool_group_descs} + assert max(len(pg.pools) for pg in native.values()) >= 2 + for g, pg in native.items(): + widths = tuple( + int(p.slot_bytes) for p in sorted(pg.pools, key=lambda p: int(p.pool_index)) + ) + assert layout.pool_groups[g] == widths + assert geometry.page_bytes[g] == sum(widths) + + # Buffers and layer groups name layers by global layer id: for virtual layers, model layer + # times the number of attention types plus a type the internal layer holds. + internal = sorted(int(i) for group in mgr.impl.layer_grouping for i in group) + global_of = dict(zip(internal, global_layer_ids(mgr, internal))) + assert len(set(global_of.values())) == len(internal) + virtual = mgr._layer_attn_to_layer_id + num_types = max(t.value for t in type(next(iter(virtual))[1])) + 1 + for (model_layer, attn_type), layer_id in virtual.items(): + assert global_of[layer_id] // num_types == model_layer + held = {t.value for (m, t), lid in virtual.items() if lid == layer_id} + assert global_of[layer_id] % num_types in held + described = {(b.layer, b.role) for b in layout.buffers} + assert described == { + (global_of[int(b.layer_id)], str(b.role)) for b in mgr.impl.all_buffer_ids + } + layers = sorted(layer for group in layout.layer_groups for layer in group.layers) + assert layers == sorted(global_of.values()) + # One KV head: every layer group is the same content on every rank. + assert all(group.shard == ShardDesc() for group in layout.layer_groups) diff --git a/tests/unittest/_torch/executor/kv_cache/sharing/test_layout_matches_native_page_table.py b/tests/unittest/_torch/executor/kv_cache/sharing/test_layout_matches_native_page_table.py new file mode 100644 index 000000000000..d47bd51bcbd6 --- /dev/null +++ b/tests/unittest/_torch/executor/kv_cache/sharing/test_layout_matches_native_page_table.py @@ -0,0 +1,313 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""``derive_layout`` against the page table native builds from the same live manager, for three +model shapes with small caches and no checkpoint: each transferred buffer's layer, pool, offset, +size and mapper, and each layer group's kind, window, heads, pools and view roles.""" + +from __future__ import annotations + +import dataclasses +import gc +from collections import defaultdict +from contextlib import contextmanager +from unittest.mock import patch + +import pytest +import torch + +from tensorrt_llm._torch.disaggregation.resource.kv_extractor import build_page_table_from_manager +from tensorrt_llm._torch.disaggregation.resource.page import CacheKind, MapperKind +from tensorrt_llm._torch.distributed.communicator import Distributed +from tensorrt_llm._torch.pyexecutor.kv_cache.sharing._layout import ( + BufferMapper, + ShardDesc, + derive_layout, +) +from tensorrt_llm.bindings import DataType +from tensorrt_llm.bindings.internal.batch_manager import CacheType +from tensorrt_llm.llmapi.llm_args import KvCacheConfig, MambaStateConfig +from tensorrt_llm.mapping import Mapping + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="allocates KV cache pools") + + +class _RanksAgree: + """The collectives of a TP2 manager built alone in this process: every rank agrees with it.""" + + local_world_size = 1 + + @staticmethod + def allreduce(value, op=None): + return value + + +@contextmanager +def _managed(factory): + gc.collect() + torch.cuda.empty_cache() + with patch.object(Distributed, "get", return_value=_RanksAgree()): + mgr = factory() + try: + yield mgr + finally: + stream = getattr(mgr, "_stream", None) + if stream is not None: + stream.synchronize() + mgr.shutdown() + del mgr + gc.collect() + torch.cuda.empty_cache() + + +def _mapping(tp=1, rank=0): + return Mapping(world_size=tp, rank=rank, tp_size=tp) + + +def _v2(kv_cache_type=CacheType.SELF, **kwargs): + from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import KVCacheManagerV2 + + return KVCacheManagerV2( + kv_cache_type=kv_cache_type, + kv_cache_config=KvCacheConfig(max_tokens=4096, enable_block_reuse=True), + max_batch_size=4, + max_seq_len=2048, + vocab_size=32000, + execution_stream=torch.cuda.Stream(), + dtype=DataType.BF16, + **kwargs, + ) + + +# TinyLlama-1.1B: 22 layers, 4 KV heads of 64. Rank 1 of TP2 holds heads 2 and 3. +def tinyllama_tp2_rank1(): + return _managed( + lambda: _v2( + num_layers=22, num_kv_heads=4, head_dim=64, tokens_per_block=32, mapping=_mapping(2, 1) + ) + ) + + +# DeepSeek-V3-Lite: 30 MLA layers, one 576-wide latent per token (512 + 64 rope), K only. +def dsv3_lite_mla(): + return _managed( + lambda: _v2( + kv_cache_type=CacheType.SELFKONLY, + num_layers=30, + num_kv_heads=1, + head_dim=512 + 64, + tokens_per_block=64, + mapping=_mapping(), + ) + ) + + +# Nemotron-Nano-9B-v2 shape, fewer layers: Mamba2 (128 heads of 80, d_state 128, 8 groups, d_conv 4) +# at layers 0, 2 and 4, attention (8 KV heads of 128) at layer 3. +def nemotron_h_hybrid(): + from tensorrt_llm._torch.pyexecutor.kv_cache.mamba_cache_manager import ( + MambaHybridCacheManagerV2, + ) + + pattern = "M-M*M-" + mamba_mask = [c == "M" for c in pattern] + attn_mask = [c == "*" for c in pattern] + return _managed( + lambda: MambaHybridCacheManagerV2( + mamba_d_state=128, + mamba_d_conv=4, + mamba_num_heads=128, + mamba_n_groups=8, + mamba_head_dim=80, + mamba_num_layers=sum(mamba_mask), + mamba_layer_mask=mamba_mask, + mamba_cache_dtype=torch.bfloat16, + mamba_ssm_cache_dtype=torch.float32, + kv_cache_config=KvCacheConfig( + max_tokens=2048, + enable_block_reuse=True, + mamba_state_config=MambaStateConfig(periodic_snapshot_interval=256), + ), + kv_cache_type=CacheType.SELF, + num_layers=sum(attn_mask), + num_kv_heads=8, + head_dim=128, + tokens_per_block=32, + max_seq_len=1024, + max_batch_size=2, + mapping=_mapping(), + layer_mask=attn_mask, + vocab_size=1024, + dtype=DataType.BF16, + ) + ) + + +# shape -> (builder, the shard every layer group names) +SHAPES = { + "tinyllama_tp2_rank1": (tinyllama_tp2_rank1, ShardDesc(2, 1)), + "dsv3_lite_mla": (dsv3_lite_mla, ShardDesc()), + "nemotron_h_hybrid": (nemotron_h_hybrid, ShardDesc()), +} + + +def _as_list(values): + return None if values is None else [int(v) for v in values] + + +def native_groups(page_table): + """What native's page table says per layer group. Pools are named by base address, layers by + the global ids of ``local_layers``; order-free parts are sorted.""" + out = [] + for lg in page_table.layer_groups: + pools = [ + (int(p.base_address), int(p.slot_bytes), int(p.num_slots)) + for p in page_table.pool_groups[lg.pool_group_idx].pools + ] + global_of = {int(ll.local_layer_id): int(ll.global_layer_id) for ll in lg.local_layers} + state = lg.kind == CacheKind.STATE + rows, roles, geometry = [], {}, {} + for pv in lg.pool_views: + key = (pools[pv.pool_idx][0], MapperKind(int(pv.mapper_kind)).name) + assert key not in roles, f"two views of pool and mapper {key}" + roles[key] = sorted(pv.pool_role) + if state: + geometry[key] = (_as_list(pv.section_bytes), pv.bytes_per_head) + for e in pv.buffer_entries: + layer = global_of[int(e["local_layer_id"])] + rows.append((layer, key[0], int(e["offset"]), int(e["size"]), key[1])) + out.append( + { + "kind": "state" if state else "attention", + "window": None if state else lg.sliding_window_size, + "heads": None if state else int(lg.kv_head_num_per_rank), + "layers": sorted(global_of.values()), + "pools": pools, + "rows": sorted(rows), + "roles": roles, + "state_geometry": geometry, + } + ) + return out + + +def layout_groups(geometry): + """The same per layer group from ``derive_layout``'s result, without local-only buffers (native + skips ignored roles); transferred buffers no group claims are listed last. Attention heads are + the head-sharded buffers' ``num_heads``: native's attention views have no per-head size.""" + layout = geometry.layout + out, claimed = [], set() + for desc in layout.layer_groups: + g = desc.pool_group + pools = [(p.base, p.slot_bytes, p.num_slots) for p in geometry.device_pools[g]] + assert layout.pool_groups[g] == tuple(w for _, w, _ in pools) + base_of = {p.index: p.base for p in geometry.device_pools[g]} + state = desc.kind == "state" + rows, roles, geometries, heads = [], defaultdict(list), defaultdict(set), set() + for i, b in enumerate(layout.buffers): + if b.layer not in desc.layers or b.pool[0] != g or not b.transfer: + continue + assert b.expansion == 1, "native reads no expansion" + claimed.add(i) + key = (base_of[b.pool[1]], b.mapper.name) + rows.append((b.layer, key[0], b.offset, b.size, key[1])) + roles[key].append(b.role) + if state: + section = b.geometry.section_bytes + section = None if section is None else tuple(section) + geometries[key].add((section, b.geometry.bytes_per_head)) + elif b.mapper in (BufferMapper.INDEXED, BufferMapper.NHD): + heads.add(b.geometry.num_heads) + assert all(len(v) == 1 for v in geometries.values()), f"roles disagree: {dict(geometries)}" + assert len(heads) <= 1, heads + out.append( + { + "kind": desc.kind, + "window": desc.window, + "heads": None if state else (heads.pop() if heads else 0), + "layers": sorted(desc.layers), + "pools": pools, + "rows": sorted(rows), + "roles": {k: sorted(set(v)) for k, v in roles.items()}, + "state_geometry": {k: (_as_list(s), h) for k, ((s, h),) in geometries.items()}, + } + ) + stray = [ + (b.layer, b.role) for i, b in enumerate(layout.buffers) if b.transfer and i not in claimed + ] + if stray: + out.append({"unclaimed buffers": stray}) + return out + + +def lies(geometry): + """Layout fields that each misstate one thing native reads.""" + layout = geometry.layout + first = next(b for b in layout.buffers if b.transfer) + + def buffers(**change): + return tuple(dataclasses.replace(b, **change) if b is first else b for b in layout.buffers) + + def first_group(**change): + return tuple( + dataclasses.replace(d, **change) if i == 0 else d + for i, d in enumerate(layout.layer_groups) + ) + + other = ( + BufferMapper.INDEXED if first.mapper is BufferMapper.REPLICATED else BufferMapper.REPLICATED + ) + out = { + "offset": {"buffers": buffers(offset=first.offset + first.size)}, + "layer": {"buffers": buffers(layer=first.layer + 1000)}, + "mapper": {"buffers": buffers(mapper=other)}, + "extra buffer": { + "buffers": layout.buffers + (dataclasses.replace(first, layer=first.layer + 1000),) + }, + "window": {"layer_groups": first_group(window=64)}, + } + # The first layer group pointed at another device pool group, when the layout has one. + mine = layout.layer_groups[0].pool_group + others = sorted(g for g in layout.pool_groups if g != mine) + if others: + out["pool group"] = {"layer_groups": first_group(pool_group=others[0])} + return out + + +@pytest.mark.parametrize("shape", list(SHAPES)) +def test_layout_matches_the_page_table_native_builds(shape): + build, shard = SHAPES[shape] + with build() as mgr: + geometry = derive_layout(mgr) + page_table = build_page_table_from_manager(mgr) + native = native_groups(page_table) + derived = layout_groups(geometry) + + layout = geometry.layout + assert layout.tokens_per_block == page_table.tokens_per_block + assert len(derived) == len(native), derived[len(native) :] + assert all(group["rows"] for group in native) + for lg, (ours, theirs) in enumerate(zip(derived, native)): + for key in theirs: + assert ours[key] == theirs[key], f"layer group {lg}: {key}" + assert [d.shard for d in layout.layer_groups] == [shard] * len(native) + assert geometry.shards == ((shard.count, shard.index),) * len(native) + # The lender finds a group's pools through either; they must agree. + assert tuple(d.pool_group for d in layout.layer_groups) == geometry.pool_group_of + assert geometry.recurrent == tuple(g["kind"] == "state" for g in native) + + # The comparison is not vacuous: each misstated field shows. + for name, change in lies(geometry).items(): + lie = dataclasses.replace(geometry, layout=dataclasses.replace(layout, **change)) + assert layout_groups(lie) != native, f"the {name} lie went unnoticed" diff --git a/tests/unittest/_torch/executor/kv_cache/sharing/test_names.py b/tests/unittest/_torch/executor/kv_cache/sharing/test_names.py new file mode 100644 index 000000000000..7c7f9804c184 --- /dev/null +++ b/tests/unittest/_torch/executor/kv_cache/sharing/test_names.py @@ -0,0 +1,187 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Row and part names: who shares a name and who never does, including the ranks of a +tensor-parallel group. One full name and one part name are pinned: stored objects stay reachable +only while the bytes stay the same.""" + +import hashlib + +import numpy as np +import pytest + +from tensorrt_llm._torch.pyexecutor.kv_cache.sharing._identity import ( + KEY_BYTES, + KEY_FORMAT_VERSION, + NAME_BYTES, + Identity, + namespace, +) + +pytestmark = pytest.mark.cpu_only + +WHOLE = (1, 0) + + +def _keys(n, salt=0): + return np.array( + [[(salt * 31 + i * 7 + b) % 256 for b in range(KEY_BYTES)] for i in range(n)], + dtype=np.uint8, + ) + + +def _identity(scope=b"", layout_id=b"layout", layers=((0,),), shards=None): + return Identity(scope, layout_id, layers, shards or [WHOLE] * len(layers)) + + +def _names(identity, layer_group=0, n=3, salt=0): + return [row.tobytes() for row in identity.names(layer_group, _keys(n, salt))] + + +def _golden(): + return Identity( + hashlib.sha256(b"golden-compute").digest(), + hashlib.sha256(b"golden-layout").digest(), + [(4, 5), (0, 1, 2, 3)], + [(4, 3), WHOLE], + ) + + +def test_a_name_is_namespace_then_key_then_the_layer_group_s_identity(): + # Local group 0 holds the largest layer, so it is canonical group 5 of six. + layers = [(50,), (0,), (10,), (20,), (30,), (40,)] + identity = _identity(layers=layers, shards=[(4, 3)] + [WHOLE] * 5) + rows = identity.names(0, _keys(3)) + assert rows.shape == (3, NAME_BYTES) and rows.dtype == np.uint8 + names = [row.tobytes() for row in rows] + keys = _keys(3) + for i, name in enumerate(names): + assert len(name) == NAME_BYTES == 54 + assert name[: len(identity.namespace)] == identity.namespace + assert name[len(identity.namespace) : -6] == keys[i].tobytes() + # canonical group 5, shard 3 of 4: big-endian uint16 each + assert name[-6:] == bytes([0, 5, 0, 4, 0, 3]) + assert len(set(names)) == 3 + + +def test_one_full_name_is_pinned(): + """Any change to these bytes must bump ``KEY_FORMAT_VERSION``, so old objects become misses.""" + key = np.frombuffer(bytes(range(KEY_BYTES)), dtype=np.uint8).reshape(1, KEY_BYTES) + assert KEY_FORMAT_VERSION == 2 + assert _golden().names(0, key)[0].tobytes().hex() == ( + "ac0cbcea5a240bde17f6bac29386876a" + "000102030405060708090a0b0c0d0e0f101112131415161718191a1b1c1d1e1f" + "000100040003" + ) + + +def test_part_names_are_pinned(): + """A part is named by the layout's first 8 bytes and the canonical groups it serves.""" + identity = _golden() + assert identity.part_name([0]) == "94ede4f3bb6ba356:lg1" + assert identity.part_name([1]) == "94ede4f3bb6ba356:lg0" + assert identity.part_name([1, 0, 0]) == "94ede4f3bb6ba356:lg0+1" + + +def test_part_names_do_not_depend_on_the_scope(): + layers = [(0,), (1,)] + a = _identity(b"model", layers=layers) + b = _identity(b"other model", layers=layers) + assert [a.part_name([g]) for g in range(2)] == [b.part_name([g]) for g in range(2)] + assert a.namespace != b.namespace + + +def test_a_replicated_group_is_one_share_of_one(): + assert _names(_identity())[0][-4:] == bytes([0, 1, 0, 0]) + + +def test_instances_that_may_exchange_bytes_compute_the_same_names(): + a = _identity(b"model", layers=[(0,), (1,)]) + b = _identity(b"model", layers=[(0,), (1,)]) + assert _names(a, 1) == _names(b, 1) + + +def test_local_group_numbering_does_not_leak_into_names(): + # The same canonical groups, numbered the other way round locally, name the same objects. + a = _identity(layers=[(0,), (1,)]) + b = _identity(layers=[(1,), (0,)]) + assert _names(a, 0) == _names(b, 1) and _names(a, 1) == _names(b, 0) + assert _names(a, 0) != _names(a, 1) + assert a.part_name([0]) == b.part_name([1]) + + +@pytest.mark.parametrize("other", [(b"model", b"other-layout"), (b"other-model", b"layout")]) +def test_a_different_scope_or_layout_never_collides(other): + a = _identity(b"model", b"layout") + b = _identity(*other) + assert not set(_names(a)) & set(_names(b)) + + +def test_namespace_fields_are_length_prefixed(): + assert namespace(b"a", b"bc") != namespace(b"ab", b"c") + + +def test_an_unknown_layer_group_or_an_empty_layout_is_rejected(): + with pytest.raises(ValueError, match="unknown layer group"): + _identity().names(1, _keys(1)) + with pytest.raises(ValueError, match="layout_id"): + _identity(b"model", b"") + with pytest.raises(ValueError, match="layout_id"): + namespace(b"model", b"") + + +def test_a_scope_longer_than_its_length_prefix_is_rejected(): + assert len(namespace(b"s" * 0xFFFF, b"layout")) == 16 + with pytest.raises(ValueError, match="at most 65535"): + namespace(b"s" * 0x10000, b"layout") + + +# -- shards -------------------------------------------------------------------------------------- + + +def test_a_replicated_group_is_named_alike_on_every_rank(): + """Every buffer of the group holds the same bytes on every rank (a single KV head, a + replicated indexer): the group is whole on each rank, so ranks share its names.""" + ranks = [_identity(b"model", layers=[(0, 1), (2, 3)]) for _ in range(4)] + assert len({tuple(_names(identity)) for identity in ranks}) == 1 + + +def test_a_head_sharded_group_is_named_per_share(): + layers = [(0, 1), (2, 3)] + ranks = [_identity(b"model", layers=layers, shards=[(4, r), (4, r)]) for r in range(4)] + names = [set(_names(identity)) for identity in ranks] + for i in range(4): + for j in range(i + 1, 4): + assert not names[i] & names[j] + + +def test_groups_of_one_rank_are_judged_one_by_one(): + """A replicated group stays shared across ranks even when another group of the same rank is + head-sharded.""" + layers = [(0, 1), (2, 3)] + a = _identity(b"m", layers=layers, shards=[(2, 0), WHOLE]) + b = _identity(b"m", layers=layers, shards=[(2, 1), WHOLE]) + assert _names(a, 1) == _names(b, 1) + assert not set(_names(a, 0)) & set(_names(b, 0)) + + +@pytest.mark.parametrize("shard", [(2, 2), (0, 0), (2, -1)]) +def test_a_shard_names_a_share_of_its_count(shard): + with pytest.raises(ValueError, match="shard"): + _identity(shards=[shard]) + + +def test_every_layer_group_has_a_shard(): + with pytest.raises(ValueError, match="shards"): + Identity(b"", b"layout", [(0,), (1,)], [WHOLE]) diff --git a/tests/unittest/_torch/executor/kv_cache/sharing/test_public_surface.py b/tests/unittest/_torch/executor/kv_cache/sharing/test_public_surface.py new file mode 100644 index 000000000000..03e4d26598d7 --- /dev/null +++ b/tests/unittest/_torch/executor/kv_cache/sharing/test_public_surface.py @@ -0,0 +1,847 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""The lender's public surface and boundaries: the eleven names, their members, what importing the +package loads, who may import its private modules, and what its public docstrings promise callers. +Every scan has a positive control that plants the fault it looks for. No GPU.""" + +import ast +import dataclasses +import importlib +import inspect +import json +import os +import re +import subprocess +import sys +import weakref +from pathlib import Path +from typing import List, Protocol, Set + +import numpy as np +import pytest + +import tensorrt_llm +from tensorrt_llm._torch.pyexecutor.kv_cache import sharing +from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import ( + GroupRun, + InPlaceLender, + Lease, + Part, + PartsHold, + Readiness, + RegionView, + StagingLender, + StagingOptions, + attach_in_place, + attach_staging, +) + +SH = "tensorrt_llm._torch.pyexecutor.kv_cache.sharing" +MGR = "tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2" +PUBLIC = [ + "GroupRun", + "InPlaceLender", + "Lease", + "Part", + "PartsHold", + "Readiness", + "RegionView", + "StagingLender", + "StagingOptions", + "attach_in_place", + "attach_staging", +] +MODULES = ["__init__.py", "_identity.py", "_layout.py", "_lender.py", "_manager.py", "_slots.py"] +MODULES = sorted(MODULES + ["_types.py"]) +PKG_DIR = Path(sharing.__file__).resolve().parent +ROOT = Path(tensorrt_llm.__file__).resolve().parent +TESTS_DIR = Path(__file__).resolve().parent +# Manager and runtime members only the manager facade may touch. +INTERNALS = { + "_reuse_token_source", + "_augment_tokens_for_block_reuse", + "_stale_block_range", + "_resize_for_connector_prefix", + "_fill_fresh_kv_pages", + "_can_publish_block_reuse", + "_stream", + "_layer_attn_to_layer_id", + "_sharing", + "host_kv_cache_block_offsets", + "kv_cache_map", +} +BACKEND_WORDS = re.compile(r"\b(kvcr|mooncake|blob|native)\b", re.IGNORECASE) + + +def fact(text: str) -> re.Pattern: + """A fact a docstring states: ``text`` matched case-insensitively, any run of whitespace for a + space, so a fact wrapped in prose or indented under a Google-style section reads the same.""" + return re.compile(r"\s+".join(re.escape(word) for word in text.split()), re.IGNORECASE) + + +# What callers rely on, per documented name: the thread rule, rank-local outcomes, copies serial +# with the forward, what keeps memory past the manager's shutdown, and the in-place preconditions. +THREADS = ( + fact("only on the manager's thread"), + fact("call no lease or lender method"), + fact("their own channel"), +) +RANK_LOCAL = (fact("this rank's own"), fact("the caller combines every rank's outcome")) +HOLD_AT_SHUTDOWN = fact( + "A hold still open at the manager's shutdown keeps the staging memory until the process exits" +) +DOC_FACTS = { + "attach_staging": (fact("block reuse on"), fact("commits no blocks")), + "attach_in_place": ( + fact("on loan at the manager's shutdown"), + fact("stay until the process exits"), + ), + "Lease": THREADS + (fact("read the view, access the memory it points to"),), + "StagingLender": THREADS + (fact("read a view, access the memory it points to"),), + "InPlaceLender": THREADS + (fact("access the lent memory"), fact("own page-table code")), + "PartsHold": ( + HOLD_AT_SHUTDOWN, + fact("release it on the manager's thread once deregistration is confirmed"), + fact("dropping it unreleased keeps the memory"), + ), + "PartsHold.release": (fact("only on the manager's thread"),), + "StagingLender.lend_read": RANK_LOCAL, + "StagingLender.lend_write": ( + fact("Fails as ``lend_read``"), + fact("no free pages"), + fact("SWA scratch reuse"), + ), + "StagingLender.readiness": ( + fact("serial with the forward"), + fact("no lender call waits for them on the CPU"), + ), + "InPlaceLender.lend_read": RANK_LOCAL + + (fact("work the manager's stream queued that still writes those pages completed"),), + "InPlaceLender.lend_write": ( + fact("all work the manager's stream queued for those pages has completed"), + fact("unscheduled, unsuspended, unshrunk and its window still"), + ), + "StagingOptions": (fact("A capacity budget, not concurrency"), fact("holes")), + "StagingLender.parts": ( + fact("an unreleased lease"), + fact("an unreleased hold"), + fact("a slot lost to a failed copy"), + ), +} +FIRST_YEAR = 2026 +COPYRIGHT = re.compile( + r"# SPDX-FileCopyrightText: Copyright \(c\) (?:(20\d\d)-)?(20\d\d) NVIDIA CORPORATION & " + r"AFFILIATES\. All rights reserved\.$" +) +LICENSE = [ + "# SPDX-License-Identifier: Apache-2.0", + "#", + '# Licensed under the Apache License, Version 2.0 (the "License");', + "# you may not use this file except in compliance with the License.", + "# You may obtain a copy of the License at", + "#", + "# http://www.apache.org/licenses/LICENSE-2.0", + "#", + "# Unless required by applicable law or agreed to in writing, software", + '# distributed under the License is distributed on an "AS IS" BASIS,', + "# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.", + "# See the License for the specific language governing permissions and", + "# limitations under the License.", +] + + +def has_header(source: str) -> bool: + """The NVIDIA header: a copyright year, or a range of years, ending in 2026 or later, then the + Apache license text.""" + lines = source.splitlines() + match = COPYRIGHT.match(lines[0]) if lines else None + if match is None or lines[1 : 1 + len(LICENSE)] != LICENSE: + return False + first, last = int(match.group(1) or match.group(2)), int(match.group(2)) + return first <= last and last >= FIRST_YEAR + + +def package_files() -> List[Path]: + return sorted(PKG_DIR.glob("*.py")) + + +def module_name(path: Path, root: Path = ROOT) -> str: + relative = path.resolve().relative_to(root.parent).with_suffix("") + parts = list(relative.parts) + if parts[-1] == "__init__": + parts.pop() + return ".".join(parts) + + +def package_of(path: Path, root: Path = ROOT) -> str: + """The package a file's relative imports resolve against.""" + name = module_name(path, root) + return name if path.name == "__init__.py" else name.rpartition(".")[0] + + +def imported_modules(source: str, package: str) -> Set[str]: + """Every module (and ``module.name``) a source imports, lazy imports and + ``importlib.import_module`` with a literal included; relative imports resolved.""" + found: Set[str] = set() + for node in ast.walk(ast.parse(source)): + if isinstance(node, ast.Import): + found.update(alias.name for alias in node.names) + elif isinstance(node, ast.ImportFrom): + if node.level: + base = package.split(".") + base = base[: len(base) - (node.level - 1)] + module = ".".join(base + ([node.module] if node.module else [])) + else: + module = node.module or "" + found.add(module) + found.update(f"{module}.{alias.name}" for alias in node.names) + elif isinstance(node, ast.Call): + func = node.func + name = func.attr if isinstance(func, ast.Attribute) else getattr(func, "id", "") + first = node.args[0] if node.args else None + if name in ("import_module", "__import__") and isinstance(first, ast.Constant): + found.add(str(first.value)) + return found + + +PRIVATE_IMPORT = re.compile(r"^" + re.escape(SH) + r"\._") + + +def private_imports(source: str, package: str) -> List[str]: + """What a source outside the package imports of the package's private modules.""" + hits = sorted(m for m in imported_modules(source, package) if PRIVATE_IMPORT.match(m)) + for node in ast.walk(ast.parse(source)): + if isinstance(node, ast.Constant) and isinstance(node.value, str): + if "kv_cache.sharing._" in node.value: + hits.append(node.value) + return hits + + +def private_importers(root: Path, pkg_dir: Path) -> dict: + """Modules under ``root`` outside ``pkg_dir`` that import the package's private modules.""" + offenders = {} + for path in root.rglob("*.py"): + if pkg_dir in path.resolve().parents: + continue + source = path.read_text(encoding="utf-8", errors="replace") + if "sharing" not in source: + continue + try: + hits = private_imports(source, package_of(path, root)) + except SyntaxError: + continue # not importable, so it imports nothing + if hits: + offenders[path.relative_to(root).as_posix()] = hits + return offenders + + +def disaggregation_imports(source: str, package: str) -> List[str]: + prefix = "tensorrt_llm._torch.disaggregation" + return sorted(m for m in imported_modules(source, package) if m.startswith(prefix)) + + +def internal_reads(source: str) -> List[str]: + """Attribute names (or ``getattr`` literals) of manager internals a source uses.""" + hits = [] + for node in ast.walk(ast.parse(source)): + if isinstance(node, ast.Attribute) and node.attr in INTERNALS: + hits.append(node.attr) + elif isinstance(node, ast.Constant) and node.value in INTERNALS: + hits.append(str(node.value)) + return sorted(hits) + + +def threads_or_finalizers(source: str) -> List[str]: + hits = [] + for node in ast.walk(ast.parse(source)): + if isinstance(node, ast.Import): + hits += [a.name for a in node.names if a.name.split(".")[0] in _THREADING] + elif isinstance(node, ast.ImportFrom) and (node.module or "").split(".")[0] in _THREADING: + hits.append(node.module) + elif isinstance(node, ast.ImportFrom) and any(a.name == "finalize" for a in node.names): + hits.append("finalize") + elif isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) and node.name == "__del__": + hits.append("__del__") + elif isinstance(node, ast.Attribute) and node.attr == "finalize": + hits.append("finalize") + return hits + + +_THREADING = {"threading", "_thread", "concurrent", "multiprocessing"} + + +def run_python(code: str) -> dict: + """Run ``code`` in a fresh interpreter importing this tree's ``tensorrt_llm``; it prints JSON.""" + env = dict(os.environ) + env["PYTHONPATH"] = os.pathsep.join(filter(None, [str(ROOT.parent), env.get("PYTHONPATH")])) + done = subprocess.run( + [sys.executable, "-c", code], + cwd=str(ROOT.parent), + env=env, + capture_output=True, + text=True, + timeout=600, + ) + assert done.returncode == 0, done.stderr[-4000:] + result = json.loads(done.stdout.strip().splitlines()[-1]) + # Proof the child imported the package under test, not another installed copy. + assert Path(result["file"]).resolve() == Path(sharing.__file__).resolve() + return result + + +# -- the names -------------------------------------------------------------------------------- + + +@pytest.mark.cpu_only +def test_the_package_exports_exactly_eleven_names(): + assert sorted(sharing.__all__) == PUBLIC + assert len(sharing.__all__) == len(set(sharing.__all__)) + for name in PUBLIC: + assert getattr(sharing, name) is not None + + +@pytest.mark.cpu_only +def test_importing_the_package_loads_its_types_alone_and_exposes_nothing_else(): + result = run_python( + "import json, sys\n" + f"import {SH} as s\n" + "print(json.dumps({'file': s.__file__, 'all': s.__all__,\n" + " 'public': sorted(n for n in vars(s) if not n.startswith('_')),\n" + f" 'loaded': sorted(m for m in sys.modules if m.startswith('{SH}'))}}))\n" + ) + assert result["public"] == sorted(result["all"]) == PUBLIC + assert result["loaded"] == [SH, f"{SH}._types"] + + +@pytest.mark.cpu_only +def test_the_package_is_its_init_and_six_private_modules(): + assert [p.name for p in package_files()] == MODULES + subpackages = [p.name for p in PKG_DIR.iterdir() if p.is_dir() and p.name != "__pycache__"] + assert subpackages == [] + + +@pytest.mark.cpu_only +def test_the_protocols_keep_their_members_and_stay_unrelated(): + def members(cls): + return {n for n in vars(cls) if not n.startswith("_")} + + assert members(Lease) == {"poll", "failure", "mark_arrived", "release"} + assert members(StagingLender) == {"parts", "hold_parts", "lend_read", "lend_write", "readiness"} + assert members(InPlaceLender) == {"lend_read", "lend_write"} + assert members(PartsHold) == {"release"} + assert StagingLender not in InPlaceLender.__mro__ + assert InPlaceLender not in StagingLender.__mro__ + assert Lease not in PartsHold.__mro__ and PartsHold not in Lease.__mro__ + for protocol in (Lease, StagingLender, InPlaceLender, PartsHold): + assert Protocol in protocol.__mro__ + assert not isinstance(object(), protocol), "checkable at run time" + assert isinstance(Lease.failure, property) and isinstance(StagingLender.parts, property) + + def params(func): + return list(inspect.signature(func).parameters) + + assert params(Lease.poll) == ["self"] + assert params(Lease.mark_arrived) == ["self", "masks"] + assert params(Lease.release) == ["self"] + assert params(PartsHold.release) == ["self"] + assert params(StagingLender.hold_parts) == ["self"] + for protocol in (StagingLender, InPlaceLender): + assert params(protocol.lend_read) == ["self", "request", "start", "end"] + assert params(protocol.lend_write) == ["self", "request", "start", "end"] + assert params(StagingLender.readiness) == ["self", "request"] + + +def public_names(obj) -> Set[str]: + return {n for n in dir(obj) if not n.startswith("_")} + + +LEASE_MEMBERS = {"poll", "failure", "mark_arrived", "release"} +STAGING_MEMBERS = {"parts", "hold_parts", "lend_read", "lend_write", "readiness"} + + +@pytest.mark.cpu_only +def test_lenders_leases_and_holds_show_only_their_protocol(): + from types import SimpleNamespace + + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + class Owner: # anything a weak reference can point to + pass + + owner = Owner() + ref = weakref.ref(owner) + layout = SimpleNamespace(pool_groups=(), windows=()) # all the constructors read + shown = { + "staging lender": (_lender.Staging(ref, layout, None, (), None, ref), STAGING_MEMBERS), + "in-place lender": (_lender.InPlace(ref, layout), {"lend_read", "lend_write"}), + "staging lease": (_lender._StagingLease(owner, "read", 1), LEASE_MEMBERS), + "in-place lease": (_lender._InPlaceLease(owner, "write", None), LEASE_MEMBERS), + "parts hold": (_lender._PartsHold(owner), {"release"}), + } + assert {what: public_names(obj) for what, (obj, _) in shown.items()} == { + what: members for what, (_, members) in shown.items() + } + # A lease showing its view as an attribute, and a lender showing a manager hook, are found. + leaky = _lender._StagingLease(owner, "read", 1) + leaky.view = None + assert public_names(leaky) - LEASE_MEMBERS == {"view"} + hooked = type("Hooked", (_lender.InPlace,), {"on_free": _lender.InPlace._on_free}) + assert public_names(hooked(ref, layout)) - {"lend_read", "lend_write"} == {"on_free"} + + +@pytest.mark.cpu_only +def test_the_types_keep_their_fields(): + def fields(cls): + return [(f.name, f.default) for f in dataclasses.fields(cls)] + + missing = dataclasses.MISSING + assert fields(StagingOptions) == [ + ("fetch_tokens", missing), + ("max_fetches", 1), + ("max_bytes", None), + ] + assert fields(Part) == [(n, missing) for n in ("name", "address", "nbytes", "slot_bytes")] + [ + ("slots", missing) + ] + assert fields(GroupRun) == [ + ("layer_group", missing), + ("ordinals", missing), + ("names", None), + ("addresses", None), + ("part", None), + ] + assert fields(RegionView) == [("runs", missing)] + for cls in (StagingOptions, Part, GroupRun, RegionView): + assert cls.__dataclass_params__.frozen, cls + assert Readiness._fields == ("usable_until", "restart_floor") + assert issubclass(Readiness, tuple) + + def beyond_fields(cls): + named = {f.name for f in dataclasses.fields(cls)} + return {n for n in dir(cls) if not n.startswith("_")} - named + + assert beyond_fields(GroupRun) == {"select"} + assert beyond_fields(RegionView) == {"num_rows", "row_masks"} + assert beyond_fields(Part) == beyond_fields(StagingOptions) == set() + assert {n for n in dir(Readiness) if not n.startswith("_")} == { + "usable_until", + "restart_floor", + "count", + "index", + } + + +@pytest.mark.cpu_only +def test_the_attach_functions_keep_their_signatures(): + kind = inspect.Parameter + + def shape(func): + return [(p.name, p.kind, p.default) for p in inspect.signature(func).parameters.values()] + + assert shape(attach_staging) == [ + ("manager", kind.POSITIONAL_OR_KEYWORD, kind.empty), + ("scope", kind.KEYWORD_ONLY, kind.empty), + ("staging", kind.KEYWORD_ONLY, kind.empty), + ] + assert shape(attach_in_place) == [("manager", kind.POSITIONAL_OR_KEYWORD, kind.empty)] + + +@pytest.mark.cpu_only +def test_a_name_is_as_long_as_the_identity_makes_it(): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _identity, _types + + assert _types._NAME_BYTES == _identity.NAME_BYTES == 54 + + +@pytest.mark.cpu_only +def test_the_manager_gains_only_a_private_class_attribute(): + from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import KVCacheManagerV2 + + assert "_sharing" in vars(KVCacheManagerV2) and KVCacheManagerV2._sharing is None + words = re.compile(r"lend|lent|loan|sharing|staging|retain") + assert [n for n in dir(KVCacheManagerV2) if not n.startswith("_") and words.search(n)] == [] + + +# -- boundaries ------------------------------------------------------------------------------- + + +@pytest.mark.cpu_only +def test_the_scans_catch_every_import_form_they_look_for(): + package = "tensorrt_llm._torch.pyexecutor.kv_cache" + caught = [ + f"from {SH}._lender import attach_staging", + f"from {SH} import _types", + f"import {SH}._slots", + "from .sharing._identity import Identity", + "from .sharing import _layout", + f"import importlib\nimportlib.import_module('{SH}._manager')", + "def later():\n from .sharing._lender import Staging\n", + ] + for source in caught: + assert private_imports(source, package), source + for source in (f"from {SH} import GroupRun", "from .sharing import attach_staging"): + assert private_imports(source, package) == [], source + assert disaggregation_imports("from ....disaggregation.resource.page import MapperKind", SH) + assert disaggregation_imports("def f():\n import tensorrt_llm._torch.disaggregation\n", SH) + assert disaggregation_imports("from .._layout import x", SH) == [] + assert internal_reads("def f(m):\n return m._stream, getattr(m, 'kv_cache_map')\n") == [ + "_stream", + "kv_cache_map", + ] + assert threads_or_finalizers("import threading\n") == ["threading"] + assert threads_or_finalizers("class A:\n def __del__(self):\n pass\n") == ["__del__"] + assert threads_or_finalizers("import weakref\nweakref.finalize(o, f)\n") == ["finalize"] + + +@pytest.fixture +def tree_with_a_private_importer(tmp_path): + """A ``tensorrt_llm`` tree whose package imports its own private module, as it may, and one + module outside it that imports a private module too, as none may.""" + root = tmp_path / "tensorrt_llm" + pkg_dir = root / "_torch/pyexecutor/kv_cache/sharing" + pkg_dir.mkdir(parents=True) + (pkg_dir / "__init__.py").write_text("from ._types import Part\n") + (pkg_dir / "_lender.py").write_text("from ._types import Part\n") + (root / "_torch/pyexecutor/public_user.py").write_text( + "from .kv_cache.sharing import attach_staging\n" + ) + (root / "_torch/pyexecutor/private_user.py").write_text( + "def build():\n from .kv_cache.sharing._lender import Staging\n" + ) + return root.resolve(), pkg_dir.resolve() + + +@pytest.mark.cpu_only +def test_the_boundary_scan_catches_a_module_importing_a_private_one(tree_with_a_private_importer): + root, pkg_dir = tree_with_a_private_importer + assert private_importers(root, pkg_dir) == { + "_torch/pyexecutor/private_user.py": [f"{SH}._lender", f"{SH}._lender.Staging"] + } + + +@pytest.mark.cpu_only +def test_nothing_outside_the_package_imports_its_private_modules(): + assert private_importers(ROOT, PKG_DIR) == {} + + +@pytest.mark.cpu_only +def test_the_package_imports_nothing_from_disaggregation(): + offenders = {} + for path in package_files(): + hits = disaggregation_imports(path.read_text(), package_of(path)) + if hits: + offenders[path.name] = hits + assert offenders == {} + + +@pytest.mark.cpu_only +def test_the_manager_never_imports_the_package(): + path = ROOT / "_torch/pyexecutor/kv_cache/kv_cache_manager_v2.py" + source = path.read_text() + loaded = imported_modules(source, package_of(path)) + assert [m for m in loaded if m.startswith(SH)] == [] + planted = source + "\nfrom .sharing import attach_in_place\n" + assert [m for m in imported_modules(planted, package_of(path)) if m.startswith(SH)] + + +@pytest.mark.cpu_only +def test_the_native_path_loads_nothing_of_the_package(): + result = run_python( + "import json, sys\n" + "import tensorrt_llm._torch.disaggregation.transceiver\n" + f"import {MGR}\n" + f"loaded = sorted(m for m in sys.modules if m.startswith('{SH}'))\n" + f"import {SH} as s\n" + "print(json.dumps({'file': s.__file__, 'loaded': loaded}))\n" + ) + assert result["loaded"] == [] + + +@pytest.mark.cpu_only +def test_only_the_manager_facade_touches_manager_internals(): + offenders = {} + for path in package_files(): + if path.name == "_manager.py": + continue + hits = internal_reads(path.read_text()) + if hits: + offenders[path.name] = hits + assert offenders == {} + + +@pytest.mark.cpu_only +def test_the_package_has_no_threads_locks_or_finalizers(): + offenders = {} + for path in package_files(): + hits = threads_or_finalizers(path.read_text()) + if hits: + offenders[path.name] = hits + assert offenders == {} + + +# -- public docstrings and headers ------------------------------------------------------------ + + +def public_docstrings(): + """(where, docstring) for the package, each exported name and each of their public members.""" + yield SH, sharing.__doc__ or "" + for name in PUBLIC: + obj = getattr(sharing, name) + yield name, obj.__doc__ or "" + if not inspect.isclass(obj): + continue + for member, value in vars(obj).items(): + if member.startswith("_"): + continue + if isinstance(value, (staticmethod, classmethod)): + value = value.__func__ + doc = getattr(value, "__doc__", None) + if doc and (inspect.isfunction(value) or isinstance(value, property)): + yield f"{name}.{member}", doc + + +def backend_words(docstrings) -> List[str]: + return [where for where, doc in docstrings if BACKEND_WORDS.search(doc)] + + +@pytest.mark.cpu_only +def test_public_docstrings_name_no_backend(monkeypatch): + documented = list(public_docstrings()) + assert {where for where, _ in documented} >= {"attach_staging", "Lease.poll", "Part"} + assert backend_words(documented) == [] + # A planted word in one member's docstring is found. + doc = StagingLender.lend_read.__doc__ + monkeypatch.setattr(StagingLender.lend_read, "__doc__", doc + " Suits a Mooncake store.") + assert backend_words(public_docstrings()) == ["StagingLender.lend_read"] + + +def missing_facts(docstrings) -> dict: + """Per documented name of ``DOC_FACTS``, the facts its docstring does not state.""" + docs = dict(docstrings) + missing = {} + for where, facts in DOC_FACTS.items(): + lacking = [f.pattern for f in facts if not f.search(docs.get(where, ""))] + if lacking: + missing[where] = lacking + return missing + + +GOOGLE_STYLE_READINESS = """Where the request may resume once its fetch settled. + + Args: + request: The request a fetch went into. + + Returns: + ``None`` while a fetch into the request is unsettled. Copies queue on the manager's stream, + serial with the forward + on the GPU; with page-locked staging no lender call waits for them on the + CPU. + """ + + +@pytest.mark.cpu_only +def test_public_docstrings_state_what_callers_rely_on(monkeypatch): + assert missing_facts(public_docstrings()) == {} + # A Google-style docstring stating the same facts across its sections passes. + monkeypatch.setattr(StagingLender.readiness, "__doc__", GOOGLE_STYLE_READINESS) + assert missing_facts(public_docstrings()) == {} + # Docstrings that drop a fact are found, and a hold whose release reads as freeing. + monkeypatch.setattr(InPlaceLender, "__doc__", "Lends a request's own device pages.") + monkeypatch.setattr(StagingLender.lend_read, "__doc__", "A copy of the committed blocks.") + monkeypatch.setattr( + StagingLender.readiness, "__doc__", GOOGLE_STYLE_READINESS.replace("serial", "parallel") + ) + hold_doc = " ".join(PartsHold.__doc__.split()) + frees = HOLD_AT_SHUTDOWN.sub("Releasing a hold frees the staging memory", hold_doc) + monkeypatch.setattr(PartsHold, "__doc__", frees) + missing = missing_facts(public_docstrings()) + assert sorted(missing) == [ + "InPlaceLender", + "PartsHold", + "StagingLender.lend_read", + "StagingLender.readiness", + ] + assert missing["PartsHold"] == [HOLD_AT_SHUTDOWN.pattern] + assert missing["StagingLender.readiness"] == [fact("serial with the forward").pattern] + + +def written_files() -> List[Path]: + return package_files() + sorted(TESTS_DIR.glob("*.py")) + + +@pytest.mark.cpu_only +def test_every_new_file_carries_the_license_header(): + assert [p.name for p in written_files() if not has_header(p.read_text())] == [] + body = "\n".join(LICENSE) + + def header(years: str) -> str: + owner = "NVIDIA CORPORATION & AFFILIATES. All rights reserved." + return f"# SPDX-FileCopyrightText: Copyright (c) {years} {owner}\n{body}" + + for years in ("2026", "2027", "2025-2026", "2026-2027", "2031"): + assert has_header(header(years)), years + for years in ("2025", "2024-2025", "2027-2026"): + assert not has_header(header(years)), years + assert not has_header("# SPDX-FileCopyrightText: Copyright (c) 2026 Someone else.\n" + body) + assert not has_header(header("2026").splitlines()[0] + "\n# SPDX-License-Identifier: MIT\n") + + +def identifiers(source: str) -> Set[str]: + """Every name the code defines or uses: variables, attributes, functions, classes, arguments + and imported names; comments and docstrings aside.""" + found: Set[str] = set() + for node in ast.walk(ast.parse(source)): + if isinstance(node, ast.Name): + found.add(node.id) + elif isinstance(node, ast.Attribute): + found.add(node.attr) + elif isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): + found.add(node.name) + elif isinstance(node, ast.arg): + found.add(node.arg) + elif isinstance(node, ast.alias): + found.add((node.asname or node.name).split(".")[-1]) + return found + + +@pytest.mark.cpu_only +def test_the_package_computes_no_compute_identity(): + """Which instance computed a block is the caller's: the lender names blocks by layout, scope and + reuse key alone.""" + offenders = {} + for path in package_files(): + hits = sorted(n for n in identifiers(path.read_text()) if "compute_id" in n) + if hits: + offenders[path.name] = hits + assert offenders == {} + planted = '"""compute_id in prose is fine."""\ndef name(key, compute_id):\n return key\n' + assert sorted(n for n in identifiers(planted) if "compute_id" in n) == ["compute_id"] + + +# -- the public types' own rules -------------------------------------------------------------- + + +def staging_run(n=3, layer_group=0, part=0): + return GroupRun( + layer_group, + np.arange(n), + np.zeros((n, 54), np.uint8), + np.arange(n, dtype=np.int64) * 4096, + part, + ) + + +@pytest.mark.cpu_only +def test_a_group_run_carries_all_three_placements_or_none(): + ordinals = np.arange(3) + names = np.zeros((3, 54), np.uint8) + addresses = np.arange(3, dtype=np.int64) + for partial in ((names, None, None), (None, addresses, 0), (names, addresses, None)): + with pytest.raises(ValueError): + GroupRun(0, ordinals, *partial) + in_place = GroupRun(1, ordinals) + assert (in_place.names, in_place.addresses, in_place.part) == (None, None, None) + assert len(in_place) == 3 and in_place.ordinals.dtype == np.int64 + + +@pytest.mark.cpu_only +def test_a_group_run_checks_its_shapes(): + ordinals = np.arange(3) + addresses = np.arange(3, dtype=np.int64) + with pytest.raises(ValueError): + GroupRun(0, np.zeros((3, 1))) + with pytest.raises(ValueError): + GroupRun(0, ordinals, np.zeros((3, 53), np.uint8), addresses, 0) + with pytest.raises(ValueError): + GroupRun(0, ordinals, np.zeros((3, 54), np.int8), addresses, 0) + with pytest.raises(ValueError): + GroupRun(0, ordinals, np.zeros((3, 54), np.uint8), addresses[:2], 0) + with pytest.raises(ValueError): + GroupRun(-1, ordinals) + with pytest.raises(ValueError): + GroupRun(0, ordinals, np.zeros((3, 54), np.uint8), addresses, -1) + with pytest.raises(TypeError): + GroupRun(True, ordinals) + with pytest.raises(TypeError): + GroupRun(0, ordinals, np.zeros((3, 54), np.uint8), addresses, 1.0) + + +@pytest.mark.cpu_only +def test_a_group_run_s_arrays_are_read_only_views(): + ordinals = np.arange(3, dtype=np.int64) + run = GroupRun(0, ordinals, np.zeros((3, 54), np.uint8), np.arange(3, dtype=np.int64), 0) + for array in (run.ordinals, run.names, run.addresses): + assert not array.flags.writeable + with pytest.raises(ValueError): + array[0] = 1 + assert ordinals.flags.writeable, "the caller's own array is left as it was" + + +@pytest.mark.cpu_only +def test_select_keeps_the_marked_rows_in_order(): + run = staging_run(4, layer_group=2, part=1) + picked = run.select(np.array([True, False, True, True])) + assert picked.ordinals.tolist() == [0, 2, 3] + assert picked.addresses.tolist() == [0, 8192, 12288] + assert picked.names.shape == (3, 54) and (picked.layer_group, picked.part) == (2, 1) + assert GroupRun(0, np.arange(2)).select(np.array([False, True])).names is None + with pytest.raises(ValueError): + run.select(np.array([True, False])) + with pytest.raises(ValueError): + run.select(np.array([1, 0, 1, 1])) + + +@pytest.mark.cpu_only +def test_a_region_view_holds_one_run_per_layer_group(): + view = RegionView([staging_run(3, 0), staging_run(2, 1, part=1)]) + assert isinstance(view.runs, tuple) and view.num_rows == 5 + masks = view.row_masks() + assert [m.shape for m in masks] == [(3,), (2,)] + assert all(m.dtype == np.bool_ and not m.any() and m.flags.writeable for m in masks) + assert all(m.all() for m in view.row_masks(True)) + assert RegionView(()).num_rows == 0 and RegionView(()).row_masks() == () + with pytest.raises(ValueError): + RegionView((staging_run(3, 0), staging_run(1, 0))) + with pytest.raises(TypeError): + RegionView((object(),)) + + +@pytest.mark.cpu_only +def test_staging_options_take_positive_integers(): + options = StagingOptions(256) + assert (options.fetch_tokens, options.max_fetches, options.max_bytes) == (256, 1, None) + assert StagingOptions(256, max_fetches=4, max_bytes=1 << 30).max_bytes == 1 << 30 + for bad in (dict(fetch_tokens=0), dict(fetch_tokens=8, max_fetches=0)): + with pytest.raises(ValueError): + StagingOptions(**bad) + with pytest.raises(ValueError): + StagingOptions(8, max_bytes=0) + for bad in ( + dict(fetch_tokens=True), + dict(fetch_tokens=8.0), + dict(fetch_tokens=8, max_bytes=1.5), + ): + with pytest.raises(TypeError): + StagingOptions(**bad) + with pytest.raises(dataclasses.FrozenInstanceError): + options.max_fetches = 2 + + +@pytest.mark.cpu_only +def test_readiness_is_the_interval_in_field_order(): + readiness = Readiness(96, 32) + assert readiness == (96, 32) and readiness.usable_until == 96 and readiness.restart_floor == 32 + assert Part("p", 4096, 8192, 4096, 2) == Part("p", 4096, 8192, 4096, 2) + assert importlib.import_module(SH).Readiness is Readiness diff --git a/tests/unittest/_torch/executor/kv_cache/sharing/test_slots.py b/tests/unittest/_torch/executor/kv_cache/sharing/test_slots.py new file mode 100644 index 000000000000..ba71b3aab246 --- /dev/null +++ b/tests/unittest/_torch/executor/kv_cache/sharing/test_slots.py @@ -0,0 +1,358 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Staging sizing and the slot queue alone, over plain layouts: no manager, no memory, no GPU.""" + +import random +from types import SimpleNamespace + +import pytest + +from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import StagingOptions +from tensorrt_llm._torch.pyexecutor.kv_cache.sharing._slots import ( + Runs, + Slots, + fetch_rows, + slot_counts, +) +from tensorrt_llm.runtime.kv_cache_manager_v2 import AttnLifeCycle + +pytestmark = pytest.mark.cpu_only + +MiB = 1 << 20 +GiB = 1 << 30 + + +def layout(tokens_per_block, groups): + """What the sizing reads of a manager's layout, for layer groups ``(window, sink blocks, pool + group, page bytes)``.""" + return SimpleNamespace( + tokens_per_block=tokens_per_block, + windows=tuple(window for window, _, _, _ in groups), + sink_blocks=tuple(sinks for _, sinks, _, _ in groups), + pool_group_of=tuple(g for _, _, g, _ in groups), + page_bytes={g: page for _, _, g, page in groups}, + ) + + +def one_fetch(lay, fetch_tokens): + return sum(rows * lay.page_bytes[g] for g, rows in fetch_rows(lay, fetch_tokens).items()) + + +# TinyLlama-1.1B, TP=1, 32 tokens per block: 22 layers of K and V, 4 heads of 64 bf16 each. +TINYLLAMA_PAGE = 22 * 2 * 4 * 64 * 2 * 32 +TINYLLAMA = layout(32, [(None, 0, 0, TINYLLAMA_PAGE)]) + +# DeepSeek-V4-Pro as its cache manager lays it out on every rank, 128 tokens per block: a +# 128-token window, compressed full attention, and the compressors' state with an 8-token window. +DSV4_PAGES = {1: 20250624, 0: 572672, 2: 39321600} +DSV4_PRO = layout( + 128, [(128, 0, 1, DSV4_PAGES[1]), (None, 0, 0, DSV4_PAGES[0]), (8, 0, 2, DSV4_PAGES[2])] +) + + +# -- sizing ----------------------------------------------------------------------------------- + + +def test_tinyllama_stages_eight_whole_prompts(): + assert fetch_rows(TINYLLAMA, 2048) == {0: 64} + assert one_fetch(TINYLLAMA, 2048) == 44 * MiB + assert slot_counts(TINYLLAMA, StagingOptions(2048, max_fetches=8)) == {0: 8 * 64} + + +def test_deepseek_v4_windows_take_their_window_and_the_full_group_the_range(): + assert fetch_rows(DSV4_PRO, 4096) == {1: 1, 0: 32, 2: 1} + assert 74 * MiB < one_fetch(DSV4_PRO, 4096) < 75 * MiB + assert slot_counts(DSV4_PRO, StagingOptions(4096, max_fetches=8)) == {1: 8, 0: 256, 2: 8} + + +def test_deepseek_v4_long_ranges_capped_at_one_fetch_keep_one_fetch(): + one = one_fetch(DSV4_PRO, 1 << 20) + assert fetch_rows(DSV4_PRO, 1 << 20) == {1: 1, 0: 8192, 2: 1} + assert 4 * GiB < one < 5 * GiB + options = StagingOptions(1 << 20, max_fetches=8, max_bytes=one) + assert slot_counts(DSV4_PRO, options) == {1: 1, 0: 8192, 2: 1} + # A cap between one and eight fetches: every group keeps its rows of one fetch and more. + cap = 960 * GiB // 8 // 4 + counts = slot_counts(DSV4_PRO, StagingOptions(1 << 20, max_fetches=8, max_bytes=cap)) + assert counts[0] >= 8192 and counts[1] >= 1 and counts[2] >= 1 + assert sum(counts[g] * DSV4_PAGES[g] for g in counts) <= cap + + +def test_short_ranges_take_no_more_than_their_blocks_even_in_a_window(): + assert fetch_rows(DSV4_PRO, 100) == {1: 1, 0: 1, 2: 1} + assert slot_counts(DSV4_PRO, StagingOptions(100)) == {1: 1, 0: 1, 2: 1} + + +def test_sink_blocks_count_with_the_window(): + lay = layout(32, [(64, 1, 0, 100), (None, 0, 1, 10)]) + assert fetch_rows(lay, 320) == {0: 1 + 2, 1: 10} + assert one_fetch(lay, 320) == 3 * 100 + 10 * 10 + assert slot_counts(lay, StagingOptions(320, max_fetches=3)) == {0: 9, 1: 30} + + +def test_layer_groups_sharing_a_pool_group_add_up(): + lay = layout(16, [(None, 0, 0, 8), (32, 0, 0, 8)]) + assert fetch_rows(lay, 160) == {0: 10 + 2} + + +def test_a_partial_block_counts_as_a_whole_one(): + assert fetch_rows(TINYLLAMA, 33) == {0: 2} + assert fetch_rows(TINYLLAMA, 1) == {0: 1} + + +@pytest.mark.parametrize("fetch_tokens", [0, -32]) +def test_fetch_tokens_must_be_positive(fetch_tokens): + with pytest.raises(ValueError): + fetch_rows(TINYLLAMA, fetch_tokens) + + +def test_max_bytes_below_one_fetch_is_rejected(): + one = one_fetch(DSV4_PRO, 4096) + with pytest.raises(ValueError, match="one fetch"): + slot_counts(DSV4_PRO, StagingOptions(4096, max_bytes=one - 1)) + assert slot_counts(DSV4_PRO, StagingOptions(4096, max_bytes=one)) == fetch_rows(DSV4_PRO, 4096) + + +def test_max_bytes_above_max_fetches_changes_nothing(): + options = StagingOptions(4096, max_fetches=2, max_bytes=100 * one_fetch(DSV4_PRO, 4096)) + assert slot_counts(DSV4_PRO, options) == {1: 2, 0: 64, 2: 2} + + +def random_layout(rng): + """A layout with up to four layer groups over up to three pool groups, and its life cycles.""" + tpb = rng.choice([4, 8, 16, 32]) + pages = {g: rng.randint(1, 1000) for g in range(rng.randint(1, 3))} + groups, life_cycles = [], [] + for _ in range(rng.randint(1, 4)): + window = rng.choice([None, rng.randint(1, 6 * tpb)]) + sink_tokens = 0 if window is None else rng.choice([0, 0, rng.randint(1, 3 * tpb)]) + life_cycle = AttnLifeCycle.make(window, sink_tokens, tpb) + g = rng.choice(sorted(pages)) + groups.append((window, int(life_cycle.num_sink_blocks), g, pages[g])) + life_cycles.append(life_cycle) + return layout(tpb, groups), life_cycles + + +@pytest.mark.parametrize("seed", range(40)) +def test_any_whole_block_range_of_at_most_fetch_tokens_fits_one_fetch(seed): + # Rows a range [start, end) needs: its blocks that a history of end still reads, as the + # runtime's own life cycle says; the longest range ending at end needs the most. + rng = random.Random(seed) + lay, life_cycles = random_layout(rng) + tpb = lay.tokens_per_block + fetch_tokens = rng.randint(1, 12 * tpb) + bound = fetch_rows(lay, fetch_tokens) + span = fetch_tokens // tpb + for end_block in range(1, 40): + end = end_block * tpb + need = dict.fromkeys(bound, 0) + for lg, life_cycle in enumerate(life_cycles): + stale = life_cycle.get_stale_range(end, tpb) + blocks = range(max(0, end_block - span), end_block) + need[lay.pool_group_of[lg]] += sum(not stale.beg <= b < stale.end for b in blocks) + assert all(need[g] <= bound[g] for g in bound), (end, need, bound) + + +@pytest.mark.parametrize("seed", range(40)) +def test_counts_are_max_fetches_fetches_or_the_cap_rounded_down(seed): + rng = random.Random(seed) + lay, _ = random_layout(rng) + fetch_tokens = rng.randint(1, 12 * lay.tokens_per_block) + rows = fetch_rows(lay, fetch_tokens) + one = one_fetch(lay, fetch_tokens) + max_fetches = rng.randint(1, 9) + assert slot_counts(lay, StagingOptions(fetch_tokens, max_fetches)) == { + g: max_fetches * rows[g] for g in rows + } + cap = rng.randint(one, max_fetches * one) + counts = slot_counts(lay, StagingOptions(fetch_tokens, max_fetches, max_bytes=cap)) + assert sum(counts[g] * lay.page_bytes[g] for g in counts) <= cap + # As many whole fetches as the cap holds hold slots at once. + whole = cap // one + assert all(whole * rows[g] <= counts[g] <= max_fetches * rows[g] for g in rows) + + +@pytest.mark.parametrize("max_bytes_in_fetches", [None, 1, 2.5]) +def test_max_fetches_ranges_hold_slots_at_once_and_one_more_waits(max_bytes_in_fetches): + rows = fetch_rows(DSV4_PRO, 4096) + max_bytes = None + at_once = 3 + if max_bytes_in_fetches is not None: + max_bytes = int(max_bytes_in_fetches * one_fetch(DSV4_PRO, 4096)) + at_once = int(max_bytes_in_fetches) + slots = Slots(slot_counts(DSV4_PRO, StagingOptions(4096, max_fetches=3, max_bytes=max_bytes))) + for _ in range(at_once): + assert slots.take(slots.ask(rows)) is not None + assert slots.take(slots.ask(rows)) is None + assert slots.num_waiting == 1 + + +# -- the slot queue --------------------------------------------------------------------------- + + +def take_now(slots, counts): + ticket = slots.ask(counts) + runs = slots.take(ticket) + assert runs is not None, f"{counts} did not fit" + return runs + + +def test_slots_of_one_lease_are_contiguous_per_group(): + slots = Slots({0: 8, 1: 8}) + assert (slots.num_slots(0), slots.free_slots(1)) == (8, 8) + runs = take_now(slots, {0: 3, 1: 5}) + assert runs.slots(0).tolist() == [0, 1, 2] + assert runs.slots(1).tolist() == [0, 1, 2, 3, 4] + assert runs.slots(7).tolist() == [] + assert runs.slots(0).dtype.name == "int64" + + +def test_a_lease_is_all_or_nothing_across_groups(): + slots = Slots({0: 8, 1: 4}) + held = take_now(slots, {1: 3}) + ticket = slots.ask({0: 2, 1: 2}) + + assert slots.take(ticket) is None + # Group 0 had room, but nothing was taken from it. + assert slots.free_slots(0) == 8 + assert slots.free_slots(1) == 1 + + slots.give(held) + assert slots.take(ticket) is not None + assert (slots.free_slots(0), slots.free_slots(1)) == (6, 2) + + +def test_contiguity_waits_for_a_run_even_when_enough_slots_are_free(): + slots = Slots({0: 6}) + a, b, c = (take_now(slots, {0: 2}) for _ in range(3)) + slots.give(a) + slots.give(c) + assert slots.free_slots(0) == 4 + + ticket = slots.ask({0: 3}) + assert slots.take(ticket) is None, "two free runs of 2 are not a run of 3" + + slots.give(b) + assert slots.take(ticket).slots(0).tolist() == [0, 1, 2] + + +def test_waiting_leases_are_served_strictly_in_order(): + slots = Slots({0: 4}) + held = take_now(slots, {0: 3}) + first = slots.ask({0: 2}) + second = slots.ask({0: 1}) + + # One slot is free and ``second`` would fit, but ``first`` asked earlier. + assert slots.take(second) is None + assert slots.take(first) is None + assert slots.num_waiting == 2 + + slots.give(held) + assert slots.take(second) is None, "still behind the first" + assert slots.take(first).slots(0).tolist() == [0, 1] + assert slots.take(second).slots(0).tolist() == [2] + assert slots.num_waiting == 0 + + +def test_cancelling_the_head_lets_the_next_lease_through(): + slots = Slots({0: 4}) + held = take_now(slots, {0: 3}) + first = slots.ask({0: 4}) + second = slots.ask({0: 1}) + assert slots.take(second) is None + + slots.cancel(first) + assert slots.take(second).slots(0).tolist() == [3] + slots.give(held) + # Cancelling an unknown or granted ticket does nothing. + slots.cancel(first) + slots.cancel(12345) + assert (slots.free_slots(0), slots.num_waiting) == (3, 0) + + +def test_a_ticket_that_is_not_waiting_is_a_key_error(): + slots = Slots({0: 8}) + with pytest.raises(KeyError): + slots.take(12345) + ticket = slots.ask({0: 1}) + assert slots.take(ticket) is not None + with pytest.raises(KeyError): + slots.take(ticket) + + +def test_frees_coalesce_into_one_run(): + slots = Slots({0: 8}) + held = [take_now(slots, {0: 2}) for _ in range(4)] + for runs in (held[1], held[3], held[0], held[2]): + slots.give(runs) + assert slots.free_slots(0) == 8 + # Only a single coalesced run can serve the whole group at once. + assert take_now(slots, {0: 8}).slots(0).tolist() == list(range(8)) + + +def test_a_double_free_is_rejected_and_changes_nothing(): + slots = Slots({0: 4, 1: 4}) + runs = take_now(slots, {0: 2, 1: 2}) + slots.give(Runs({1: (0, 2)})) + assert (slots.free_slots(0), slots.free_slots(1)) == (2, 4) + + with pytest.raises(ValueError, match="freed twice"): + slots.give(runs) + # Group 0's run was not returned on the way to finding group 1's double free. + assert (slots.free_slots(0), slots.free_slots(1)) == (2, 4) + + +@pytest.mark.parametrize("run", [(3, 2), (-1, 1)], ids=["past_the_end", "negative"]) +def test_freeing_slots_outside_the_group_is_rejected(run): + slots = Slots({0: 4}) + take_now(slots, {0: 4}) + with pytest.raises(ValueError, match="outside"): + slots.give(Runs({0: run})) + assert slots.free_slots(0) == 0 + + +def test_freeing_slots_of_an_unknown_group_is_rejected(): + slots = Slots({0: 4}) + take_now(slots, {0: 4}) + with pytest.raises(ValueError, match="pool group 5"): + slots.give(Runs({0: (0, 4), 5: (0, 1)})) + assert slots.free_slots(0) == 0 + + +@pytest.mark.parametrize( + "counts", [{0: 5}, {3: 1}, {0: -1}], ids=["too_many", "unknown", "negative"] +) +def test_a_request_that_can_never_be_granted_is_rejected_at_once(counts): + slots = Slots({0: 4}) + with pytest.raises(ValueError): + slots.check(counts) + with pytest.raises(ValueError): + slots.ask(counts) + assert slots.num_waiting == 0 + assert slots.free_slots(0) == 4 + + +def test_check_queues_nothing(): + slots = Slots({0: 4}) + slots.check({0: 4}) + assert slots.num_waiting == 0 + # A later ask is first in line. + assert take_now(slots, {0: 4}).slots(0).tolist() == [0, 1, 2, 3] + + +def test_zero_counts_need_no_slots(): + slots = Slots({0: 1}) + take_now(slots, {0: 1}) + assert take_now(slots, {0: 0, 9: 0}).runs == {} diff --git a/tests/unittest/_torch/executor/kv_cache/sharing/test_staging_lender.py b/tests/unittest/_torch/executor/kv_cache/sharing/test_staging_lender.py new file mode 100644 index 000000000000..799ac022a81d --- /dev/null +++ b/tests/unittest/_torch/executor/kv_cache/sharing/test_staging_lender.py @@ -0,0 +1,2594 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""The staging lender's contract over real KV cache managers, through the public API alone. A rule +a lender could break unnoticed has its check run twice: against the real lender, and against a +subclass breaking exactly that rule, which the check must catch. Allocates device pools.""" + +import gc +import mmap +import threading +import weakref +from types import SimpleNamespace +from typing import List, Set + +import numpy as np +import pytest +import torch +from utils.util import skip_pre_blackwell + +from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import ( + InPlaceLender, + Lease, + PartsHold, + Readiness, + RegionView, + StagingLender, + StagingOptions, + attach_in_place, + attach_staging, +) +from tensorrt_llm._utils import prefer_pinned + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="allocates KV cache pools") + +TPB = 32 +WINDOW = 64 +PROMPT = list(range(1000, 1097)) # three whole blocks and one token +OTHER_PROMPT = list(range(5000, 5097)) +END = 96 +BLOCKS = END // TPB +WINDOWED_PROMPT = list(range(2000, 2161)) # five whole blocks and one token +OTHER_WINDOWED_PROMPT = list(range(6000, 6161)) +WINDOWED_END = 160 +SOURCE, TARGET, OTHER = 1, 2, 3 +SCOPE = b"lender-contract-suite" +# What a check raises when the lender under it breaks the rule it checks; a timeout or any other +# error fails the liar test instead. +CAUGHT = (AssertionError,) + + +def attach(mgr, *, fetch_tokens, max_fetches=1, max_bytes=None, scope=SCOPE): + """The public attach, as an integrator makes it.""" + options = StagingOptions(fetch_tokens, max_fetches, max_bytes) + return attach_staging(mgr, scope=scope, staging=options) + + +def attach_breaking(rules): + """An attach installing a lender whose ``rules`` (method name -> function, or a function of the + manager returning them) replace the real ones: a lender breaking exactly those rules.""" + + def attach_rule_breaker(mgr, *, fetch_tokens, max_fetches=1, max_bytes=None, scope=SCOPE): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + chosen = rules(mgr) if callable(rules) else rules + cls = type("RuleBreaker", (_lender.Staging,), dict(chosen)) + options = StagingOptions(fetch_tokens, max_fetches, max_bytes) + return _lender._attach_staging(mgr, scope=scope, staging=options, cls=cls) + + return attach_rule_breaker + + +def real(name): + """The real lender's rule ``name``, for a rule breaker to call around.""" + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + return getattr(_lender.Staging, name) + + +def ready_view(lease, mgr, tries=3): + """The lease's view once the copies queued so far have run; fails if it never comes.""" + for _ in range(tries): + mgr._stream.synchronize() + view = lease.poll() + if view is not None: + return view + raise AssertionError(f"the lease never became ready: {lease.failure}") + + +def by_group(view): + return {run.layer_group: run.ordinals.tolist() for run in view.runs} + + +def one_fetch_bytes(kit, mgr, fetch_tokens): + """Bytes one fetch of ``fetch_tokens`` tokens needs, from the device pages: every block for full + attention, a window's blocks for a window (no sinks here).""" + dev = kit.DevicePages(mgr) + total = 0 + for lg, window in enumerate(kit.windows(mgr)): + blocks = -(-fetch_tokens // TPB) + if window is not None: + blocks = min(blocks, -(-window // TPB)) + total += blocks * dev.page_bytes(lg) + return total + + +def check_staged(kit, mgr, lender, request, view): + """Every row's slot, inside its part, holds the request's device page of that block.""" + kv = kit.kv(mgr, request) + dev = kit.DevicePages(mgr) + rows = 0 + for run in view.runs: + part = lender.parts[run.part] + assert part.slot_bytes == dev.page_bytes(run.layer_group) + slots = kit.pages(kv, run.layer_group) + for address, ordinal in zip(run.addresses.tolist(), run.ordinals.tolist()): + assert part.address <= address < part.address + part.nbytes + assert (address - part.address) % part.slot_bytes == 0 + device = kit.digest(dev.read(run.layer_group, slots[ordinal])) + assert kit.digest(kit.host_bytes(address, part.slot_bytes)) == device + rows += 1 + return rows + + +def state_of(kit, mgr, lender, request): + """Capacity, history, committed tokens and readiness of the request, or what stands for none.""" + kv = kit.kv(mgr, request) + try: + readiness = lender.readiness(request) + except ValueError: + readiness = "no cache" + if kv is None: + return None, readiness + return (kv.capacity, kv.history_length, kv.num_committed_tokens), readiness + + +# -- attach ----------------------------------------------------------------------------------- + + +def test_attach_takes_only_a_v2_manager(): + with pytest.raises(TypeError, match="KVCacheManagerV2"): + attach(SimpleNamespace(), fetch_tokens=END) + with pytest.raises(TypeError, match="KVCacheManagerV2"): + attach_in_place(SimpleNamespace()) + + +def test_a_manager_takes_one_lender_for_its_life(real_manager): + with real_manager() as mgr: + with pytest.raises(TypeError): + attach_staging(mgr, scope="text", staging=StagingOptions(END)) + lender = attach(mgr, fetch_tokens=END) # the refused call attached nothing + assert isinstance(lender, StagingLender) + with pytest.raises(ValueError, match="already attached"): + attach(mgr, fetch_tokens=END) + with pytest.raises(ValueError, match="already attached"): + attach_in_place(mgr) + + +def test_recurrent_state_is_refused(kit, hybrid_manager): + before = len(kit.retained()) + with hybrid_manager() as mgr: + with pytest.raises(ValueError, match="recurrent"): + attach(mgr, fetch_tokens=256) + with pytest.raises(ValueError, match="recurrent"): + attach_in_place(mgr) + assert len(kit.retained()) == before, "a refused attach allocated staging" + + +def stand_in(cp_type=None): + """The smallest manager the attach path reads before any other work: an uninitialised + ``KVCacheManagerV2`` whose mapping is rank 0 of two ``cp_type`` ranks, or has no mapping.""" + from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import KVCacheManagerV2 + from tensorrt_llm.mapping import Mapping + + manager = KVCacheManagerV2.__new__(KVCacheManagerV2) + if cp_type is not None: + manager.mapping = Mapping(world_size=2, cp_size=2, rank=0, cp_config={"cp_type": cp_type}) + return manager + + +def refusal(call, kind=ValueError) -> str: + """The message of the ``kind`` error ``call`` raises; the check fails on anything else.""" + try: + call() + except kind as error: + return str(error) + except AssertionError: + raise + except Exception as error: # handed to the check as its failure + raise AssertionError(f"expected a {kind.__name__}, got {error!r}") from error + raise AssertionError(NOT_REFUSED) + + +PAST_THE_CHECKS = "the attach got past its checks" +NOT_REFUSED = "the call went through" + + +def check_refused_before_any_other_work(monkeypatch, cases): + """Each ``(manager, kind, words)`` is refused by both attaches with a ``kind`` error naming + ``words``, before the layout is derived: a tripwire there fails the check.""" + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + def tripwire(manager): + raise AssertionError(PAST_THE_CHECKS) + + monkeypatch.setattr(_lender, "derive_layout", tripwire) + before = len(_lender._retained()) + for make, kind, words in cases: + staging, in_place = make(), make() + assert words in refusal(lambda: attach(staging, fetch_tokens=END), kind) + assert words in refusal(lambda: attach_in_place(in_place), kind) + assert "_sharing" not in vars(staging) and "_sharing" not in vars(in_place) + assert len(_lender._retained()) == before, "a refused attach allocated staging" + + +CONTEXT_PARALLEL = [ + (lambda: stand_in("HELIX"), ValueError, "context parallelism"), + (lambda: stand_in("ULYSSES"), ValueError, "context parallelism"), +] +NO_MAPPING = [(stand_in, TypeError, "mapping")] + + +def test_context_parallelism_is_refused(monkeypatch): + check_refused_before_any_other_work(monkeypatch, CONTEXT_PARALLEL) + + +def test_the_check_catches_a_lender_accepting_context_parallelism(monkeypatch): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + monkeypatch.setattr(_lender, "_context_parallel_size", lambda manager: 1) + with pytest.raises(CAUGHT, match=PAST_THE_CHECKS): + check_refused_before_any_other_work(monkeypatch, CONTEXT_PARALLEL) + + +def test_a_manager_without_a_mapping_is_refused(monkeypatch): + check_refused_before_any_other_work(monkeypatch, NO_MAPPING) + + +def test_the_check_catches_a_lender_defaulting_a_missing_mapping(monkeypatch): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + def defaulting(manager): + return int(getattr(getattr(manager, "mapping", None), "cp_size", 1)) + + monkeypatch.setattr(_lender, "_context_parallel_size", defaulting) + with pytest.raises(CAUGHT, match=PAST_THE_CHECKS): + check_refused_before_any_other_work(monkeypatch, NO_MAPPING) + + +def unpaired_draft(mgr) -> None: + """The state of a draft manager without joint reuse: it never publishes blocks for reuse.""" + mgr._can_publish_block_reuse = False + + +COMMITS_NOTHING = { + "block_reuse_off": ({"enable_block_reuse": False}, None), + "unpaired_draft": ({}, unpaired_draft), +} + + +def check_staging_refuses_a_manager_that_commits_nothing(kit, real_manager, attach, case): + """Staging publishes committed blocks only, so a manager that commits none is refused at the + attach, which then changed nothing; the in-place lender, which lends pages, attaches.""" + manager_kwargs, adjust = COMMITS_NOTHING[case] + with real_manager(**manager_kwargs) as mgr: + if adjust is not None: + adjust(mgr) + before = len(kit.retained()) + message = refusal(lambda: attach(mgr, fetch_tokens=END)) + assert "block reuse" in message, message + assert "_sharing" not in vars(mgr) and len(kit.retained()) == before + assert isinstance(attach_in_place(mgr), InPlaceLender) + + +@pytest.mark.parametrize("case", list(COMMITS_NOTHING)) +def test_staging_refuses_a_manager_that_commits_no_blocks(kit, real_manager, case): + check_staging_refuses_a_manager_that_commits_nothing(kit, real_manager, attach, case) + + +@pytest.mark.parametrize("case", list(COMMITS_NOTHING)) +def test_the_check_catches_a_staging_attach_ignoring_block_reuse( + kit, real_manager, monkeypatch, case +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + monkeypatch.setattr(_lender, "_check_commits", lambda manager: None) + with pytest.raises(CAUGHT, match=NOT_REFUSED): + check_staging_refuses_a_manager_that_commits_nothing(kit, real_manager, attach, case) + + +# -- sizing ----------------------------------------------------------------------------------- + + +def test_staging_holds_max_fetches_fetches_at_once(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + one = one_fetch_bytes(kit, mgr, END) + lender = attach(mgr, fetch_tokens=END, max_fetches=2) + (part,) = lender.parts + assert part.slot_bytes == kit.DevicePages(mgr).page_bytes(0) + assert part.nbytes == part.slots * part.slot_bytes == 2 * one + leases = [lender.lend_read(source, 0, END) for _ in range(3)] + for lease in leases[:2]: + ready_view(lease, mgr) + assert leases[2].poll() is None and leases[2].failure is None, "three fetches in two" + leases[0].release() + check_staged(kit, mgr, lender, source, ready_view(leases[2], mgr)) + for lease in leases[1:]: + lease.release() + + +@pytest.mark.parametrize("cap", ["one_fetch", "two_and_a_half_fetches"]) +def test_max_bytes_caps_staging_but_never_below_one_fetch(kit, real_manager, cap): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + one = one_fetch_bytes(kit, mgr, END) + with pytest.raises(ValueError, match="one fetch"): + attach(mgr, fetch_tokens=END, max_fetches=4, max_bytes=one - 1) + max_bytes, holding = (one, 1) if cap == "one_fetch" else (2 * one + one // 2, 2) + lender = attach(mgr, fetch_tokens=END, max_fetches=4, max_bytes=max_bytes) + (part,) = lender.parts + assert holding * one <= part.nbytes <= max_bytes + leases = [lender.lend_read(source, 0, END) for _ in range(holding + 1)] + for lease in leases[:holding]: + ready_view(lease, mgr) + assert leases[-1].poll() is None and leases[-1].failure is None + for lease in leases: + lease.release() + + +@pytest.mark.parametrize("num_layers", [2, 3], ids=["one_pool_group", "two_pool_groups"]) +def test_a_windowed_range_of_fetch_tokens_always_fits(kit, real_manager, num_layers): + with real_manager(windows=[WINDOW, 256], num_layers=num_layers) as mgr: + source = kit.published(mgr, SOURCE, WINDOWED_PROMPT) + target = kit.admitted(mgr, TARGET, OTHER_WINDOWED_PROMPT) + lender = attach(mgr, fetch_tokens=WINDOWED_END) + dev = kit.DevicePages(mgr) + rows = {} + for lg, (window, g) in enumerate(zip(kit.windows(mgr), kit.pool_group_of(mgr))): + rows[g] = rows.get(g, 0) + (WINDOW // TPB if window else WINDOWED_END // TPB) + groups = kit.pool_group_ids(mgr) + assert len(lender.parts) == len(groups) == num_layers - 1 + for g, part in zip(groups, lender.parts): + assert (part.slots, part.slot_bytes) == (rows[g], dev.group_page_bytes(g)) + assert part.nbytes == part.slots * part.slot_bytes + read = lender.lend_read(source, 0, WINDOWED_END) # takes every slot + assert check_staged(kit, mgr, lender, source, ready_view(read, mgr)) == sum(rows.values()) + read.release() + write = lender.lend_write(target, 0, WINDOWED_END) + view = ready_view(write, mgr) + assert view.num_rows == sum(rows.values()) + write.mark_arrived(view.row_masks()) + write.release() + + +def check_the_staging_memory_is_the_parts_and_nothing_more(kit, real_manager, attach): + # Three layers: two pool groups, so two parts. + with real_manager(windows=[WINDOW, 256], num_layers=3) as mgr: + lender = attach(mgr, fetch_tokens=WINDOWED_END, max_fetches=2) + parts = lender.parts + assert len(parts) == len(kit.pool_group_ids(mgr)) == 2 + (memory,) = kit.staging_memory(parts) + assert memory.address <= parts[0].address < memory.address + memory.nbytes + assert memory.nbytes == sum(p.nbytes for p in parts), "staging is not the sum of its parts" + + +def test_the_staging_memory_is_the_parts_and_nothing_more(kit, real_manager): + check_the_staging_memory_is_the_parts_and_nothing_more(kit, real_manager, attach) + + +def test_the_check_catches_a_lender_allocating_more_than_its_parts(kit, real_manager, monkeypatch): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + allocate = _lender._allocate + monkeypatch.setattr(_lender, "_allocate", lambda nbytes: allocate(nbytes + 4096)) + with pytest.raises(CAUGHT, match="not the sum of its parts"): + check_the_staging_memory_is_the_parts_and_nothing_more(kit, real_manager, attach) + + +class FreeSpy: + """The CUDA runtime, recording the address of every page-locked host allocation freed.""" + + def __init__(self, runtime): + self._runtime = runtime + self.freed = [] + + def __getattr__(self, name): + found = getattr(self._runtime, name) + if name != "cudaFreeHost": + return found + + def free(address): + self.freed.append(int(address)) + return found(address) + + return free + + +def check_the_staging_memory_is_pinned_at_its_size_and_freed_once( + kit, real_manager, attach, monkeypatch +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + spy = FreeSpy(_lender.cudart) + monkeypatch.setattr(_lender, "cudart", spy) + with real_manager() as mgr: + parts = attach(mgr, fetch_tokens=END, max_fetches=3).parts + base, total = parts[0].address, sum(p.nbytes for p in parts) + whole_pages = -(-total // mmap.PAGESIZE) * mmap.PAGESIZE + assert 1 << (total - 1).bit_length() > whole_pages, "a power of two would pass unseen" + start, pinned = kit.pinned_range(base) + assert start == base and total <= pinned <= whole_pages, ( + f"{pinned} bytes pinned for {total}" + ) + mgr.shutdown() + assert spy.freed == [base] and kit.pinned_range(base) is None, "not freed at the shutdown" + mgr.shutdown() + gc.collect() + assert spy.freed == [base], "freed more than once" + + +def test_the_staging_memory_is_pinned_at_its_size_and_freed_once(kit, real_manager, monkeypatch): + if not prefer_pinned(): + pytest.skip("staging is pageable where pinning does not pay off") + check_the_staging_memory_is_pinned_at_its_size_and_freed_once( + kit, real_manager, attach, monkeypatch + ) + + +def test_the_check_catches_a_lender_pinning_through_torch_s_rounding_allocator( + kit, real_manager, monkeypatch +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + class ThroughTorch(_lender._HostMemory): + def __init__(self, nbytes): + self.nbytes = nbytes + self._tensor = torch.empty(nbytes, dtype=torch.uint8, pin_memory=True) + self.address = self._tensor.data_ptr() + + def free(self): + self.address, self._tensor = 0, None + + monkeypatch.setattr(_lender, "_HostMemory", ThroughTorch) + with pytest.raises(CAUGHT, match="bytes pinned for"): + check_the_staging_memory_is_pinned_at_its_size_and_freed_once( + kit, real_manager, attach, monkeypatch + ) + + +def test_the_check_catches_a_lender_never_freeing_the_staging_memory( + kit, real_manager, monkeypatch +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + if not prefer_pinned(): + pytest.skip("staging is pageable where pinning does not pay off") + monkeypatch.setattr(_lender._HostMemory, "free", lambda self: None) + with pytest.raises(CAUGHT, match="not freed at the shutdown"): + check_the_staging_memory_is_pinned_at_its_size_and_freed_once( + kit, real_manager, attach, monkeypatch + ) + + +# -- publish ---------------------------------------------------------------------------------- + + +def test_a_publish_lends_the_committed_pages_under_their_reuse_keys(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=END) + lease = lender.lend_read(source, 0, END) + assert isinstance(lease, Lease) + view = ready_view(lease, mgr) + assert isinstance(view, RegionView) and lease.poll() is view, "one view, every poll" + (run,) = view.runs + assert run.ordinals.tolist() == list(range(BLOCKS)) and run.part == 0 + assert check_staged(kit, mgr, lender, source, view) == BLOCKS + kv = kit.kv(mgr, source) + assert [bytes(name[16:48]) for name in run.names] == kit.chain_keys(kv, PROMPT)[:BLOCKS] + assert len({bytes(name[:16]) for name in run.names}) == 1, "one namespace" + assert len({bytes(name[48:]) for name in run.names}) == 1, "one layer group, one shard" + # The manager committed under those keys: its tree finds the block after them by them. + extra = list(range(7000, 7000 + TPB)) + probed = mgr.impl.probe_first_new_block_key(kv.reuse_scope, PROMPT[:END] + extra) + assert probed == kit.chain_keys(kv, PROMPT[:END] + extra)[BLOCKS] + lease.release() + with pytest.raises(RuntimeError): + lease.poll() + + +def test_a_prompt_ending_on_a_block_boundary_publishes_its_last_block(kit, real_manager): + with real_manager() as mgr: + tokens = PROMPT[:END] + source = kit.published(mgr, SOURCE, tokens) + lender = attach(mgr, fetch_tokens=END) + lease = lender.lend_read(source, 0, END) + (run,) = ready_view(lease, mgr).runs + keys = kit.chain_keys(kit.kv(mgr, source), tokens) + assert len(keys) == BLOCKS and [bytes(n[16:48]) for n in run.names] == keys + lease.release() + + +def test_a_multimodal_publish_is_named_by_the_keys_the_manager_commits(kit, real_manager): + image = dict(multimodal_positions=[40], multimodal_lengths=[16]) + with real_manager() as mgr: + first = kit.published(mgr, SOURCE, PROMPT, multimodal_hashes=[list(range(1, 9))], **image) + second = kit.published( + mgr, OTHER, PROMPT, multimodal_hashes=[list(range(8, 0, -1))], **image + ) + lender = attach(mgr, fetch_tokens=END, max_fetches=2) + leases = [lender.lend_read(r, 0, END) for r in (first, second)] + names = [ready_view(lease, mgr).runs[0].names for lease in leases] + keys = [[bytes(n[16:48]) for n in run_names] for run_names in names] + kv = kit.kv(mgr, first) + augmented = list(mgr._augment_tokens_for_block_reuse(PROMPT, first)) + assert keys[0] == kit.chain_keys(kv, augmented)[:BLOCKS] + extra = list(range(7000, 7000 + TPB)) + probed = mgr.impl.probe_first_new_block_key(kv.reuse_scope, augmented[:END] + extra) + assert probed == kit.chain_keys(kv, augmented[:END] + extra)[BLOCKS] + # The image sits in block 1: the block before it is shared, the rest is not. + assert keys[0][0] == keys[1][0] and not set(keys[0][1:]) & set(keys[1][1:]) + assert keys[0][1:] != kit.chain_keys(kv, PROMPT)[1:BLOCKS], "the image changes the keys" + for lease in leases: + lease.release() + + +def test_a_windowed_publish_lends_only_what_the_window_at_its_end_reads(kit, real_manager): + with real_manager(windows=[WINDOW, 256]) as mgr: + source = kit.published(mgr, SOURCE, WINDOWED_PROMPT) + lender = attach(mgr, fetch_tokens=WINDOWED_END) + windows = kit.windows(mgr) + sliding, full = windows.index(WINDOW), windows.index(None) + lease = lender.lend_read(source, 0, WINDOWED_END) + view = ready_view(lease, mgr) + beg, end = kit.stale_blocks(mgr, sliding, WINDOWED_END) + assert end > beg, "the window must leave early blocks behind" + blocks = range(WINDOWED_END // TPB) + in_window = [o for o in blocks if not beg <= o < end] + assert by_group(view) == {full: list(blocks), sliding: in_window} + addresses = np.concatenate([run.addresses for run in view.runs]) + assert len(set(addresses.tolist())) == len(addresses), "every row has its own slot" + assert check_staged(kit, mgr, lender, source, view) == len(addresses) + lease.release() + + # The request's own window has since passed what a window at 96 reads: none is lent. + early = lender.lend_read(source, 0, END) + rows = by_group(ready_view(early, mgr)) + history = kit.kv(mgr, source).history_length + passed_beg, passed_end = kit.stale_blocks(mgr, sliding, history) + assert rows[full] == list(range(BLOCKS)) + assert [o for o in rows.get(sliding, []) if passed_beg <= o < passed_end] == [] + assert rows.get(sliding, []) == [] + early.release() + + +def test_a_range_without_rows_is_ready_at_once_and_waits_for_no_one(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=END) + holder = lender.lend_read(source, 0, END) # every slot + ready_view(holder, mgr) + empty = lender.lend_read(source, END, END) + view = empty.poll() + assert view is not None and view.num_rows == 0 + assert len(view.runs) == kit.num_layer_groups(mgr) + empty.release() + holder.release() + + +# -- ready means copied ----------------------------------------------------------------------- + + +def check_ready_means_copied(kit, real_manager, attach): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=END) + with kit.held_stream(mgr._stream) as gate: + lease = lender.lend_read(source, 0, END) + assert lease.poll() is None, "ready before its copy into staging ran" + assert lease.poll() is None and lease.failure is None + gate.open() + check_staged(kit, mgr, lender, source, ready_view(lease, mgr)) + lease.release() + + +def test_a_publish_is_ready_only_once_its_copy_ran(kit, real_manager): + check_ready_means_copied(kit, real_manager, attach) + + +def test_the_check_catches_a_lender_ready_before_its_copy(kit, real_manager): + liar = attach_breaking({"_copy_landed": lambda self, lease: True}) + with pytest.raises(CAUGHT, match="ready before its copy into staging ran"): + check_ready_means_copied(kit, real_manager, liar) + + +# -- no slot reused while a copy or a backend may touch it ----------------------------------- + + +def check_slots_are_not_reused_while_touched(kit, real_manager, attach): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + first = kit.admitted(mgr, TARGET, OTHER_PROMPT) + second = kit.admitted(mgr, OTHER, OTHER_WINDOWED_PROMPT) + for request in (first, second): # grown ahead, so the held section only lends + assert mgr._resize_for_connector_prefix(request, kit.kv(mgr, request), 0, END) + lender = attach(mgr, fetch_tokens=END) # one fetch: three slots + with kit.held_stream(mgr._stream) as gate: + read = lender.lend_read(source, 0, END) # every slot, its copy held + read.release() + write = lender.lend_write(first, 0, END) + assert write.poll() is None, "slots handed out while a copy into them was queued" + gate.open() + view = ready_view(write, mgr) + write.release() # seen ready and not marked: its backend may still write the slots + waiting = lender.lend_write(second, 0, END) + assert waiting.poll() is None, "slots handed out while a backend could still write them" + write.mark_arrived(view.row_masks()) + waiting_view = ready_view(waiting, mgr) + waiting.mark_arrived(waiting_view.row_masks()) + waiting.release() + + +def test_a_slot_returns_only_after_its_copy_and_its_marks(kit, real_manager): + check_slots_are_not_reused_while_touched(kit, real_manager, attach) + + +def test_the_check_catches_a_lender_recycling_before_the_copy_lands(kit, real_manager): + def recycle_early(self, lease): + with kit.events_report_done(): + return real("_recyclable")(self, lease) + + with pytest.raises(CAUGHT, match="slots handed out while a copy into them was queued"): + check_slots_are_not_reused_while_touched( + kit, real_manager, attach_breaking({"_recyclable": recycle_early}) + ) + + +def test_leases_wait_for_slots_strictly_in_order(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=END) # three slots + first = lender.lend_read(source, 0, 2 * TPB) + ready_view(first, mgr) + second = lender.lend_read(source, 0, END) + third = lender.lend_read(source, 0, TPB) + assert second.poll() is None + assert third.poll() is None, "a later lease that fits jumped the one waiting ahead" + first.release() + ready_view(second, mgr) + assert third.poll() is None + second.release() + assert by_group(ready_view(third, mgr)) == {0: [0]} + third.release() + + +def test_max_fetches_is_a_budget_and_holes_can_make_a_lease_wait(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=END, max_fetches=2) + (group,) = kit.pool_group_ids(mgr) + (part,) = lender.parts + assert part.slots == 2 * BLOCKS + first = lender.lend_read(source, 0, 2 * TPB) # slots 0 and 1 + second = lender.lend_read(source, 0, TPB) # slot 2 + third = lender.lend_read(source, 0, TPB) # slot 3 + for lease in (first, second, third): + ready_view(lease, mgr) + first.release() + assert lender._free_slots(group) == 4 # slots 0, 1, 4 and 5 + fetch = lender.lend_read(source, 0, END) + assert fetch.poll() is None and fetch.failure is None, "no three free slots in a row" + second.release() # slots 0 to 2 are free in a row + (run,) = ready_view(fetch, mgr).runs + assert ((run.addresses - part.address) // part.slot_bytes).tolist() == [0, 1, 2] + for lease in (third, fetch): + lease.release() + + +# -- fetch ------------------------------------------------------------------------------------ + + +def check_a_fetch_settles_when_its_copy_lands(kit, real_manager, attach, local_blocks=0): + with real_manager() as mgr_a, real_manager() as mgr_b: + source = kit.published(mgr_a, SOURCE, PROMPT) + lender_a = attach(mgr_a, fetch_tokens=END) + publish = lender_a.lend_read(source, 0, END) + publish_view = ready_view(publish, mgr_a) + if local_blocks: + kit.published(mgr_b, 7, PROMPT[: local_blocks * TPB] + list(range(9000, 9033))) + target = kit.admitted(mgr_b, TARGET, PROMPT) + kv = kit.kv(mgr_b, target) + local, history = kv.num_committed_tokens, kv.history_length + assert local == local_blocks * TPB + lender_b = attach(mgr_b, fetch_tokens=END) + assert lender_b.readiness(target) == (max(local, history), history) + + lease = lender_b.lend_write(target, local, END) + assert lender_b.readiness(target) is None, "a fetch counts from its call" + assert kv.capacity >= END and kv.history_length == history + view = lease.poll() + assert view is not None, "a write is ready at its first poll after the grant" + assert by_group(view) == {0: list(range(local_blocks, BLOCKS))} + kit.fill_sentinel(mgr_b, target, first_block=local_blocks) + masks = kit.relay(lender_a, publish_view, lender_b, view) + with kit.held_stream(mgr_b._stream) as gate: + lease.mark_arrived(masks) + assert lender_b.readiness(target) is None, "settled before the copy into pages ran" + gate.open() + readiness = lender_b.readiness(target) + assert isinstance(readiness, Readiness) and all(type(v) is int for v in readiness) + assert readiness == (END, min(local, history)) + for ordinal in range(BLOCKS): + got = kit.digest(kit.page(mgr_b, target, 0, ordinal)) + assert got == kit.digest(kit.page(mgr_a, source, 0, ordinal)), "a fetched page differs" + lease.release() + publish.release() + + +@pytest.mark.parametrize("local_blocks", [0, 2], ids=["cold", "local_prefix"]) +def test_a_fetch_is_usable_once_its_copy_lands(kit, real_manager, local_blocks): + check_a_fetch_settles_when_its_copy_lands(kit, real_manager, attach, local_blocks) + + +def test_the_check_catches_readiness_before_the_copy(kit, real_manager): + def settled_early(self, fetch): + with kit.events_report_done(): + return real("_report_settled")(self, fetch) + + with pytest.raises(CAUGHT, match="settled before the copy into pages ran"): + check_a_fetch_settles_when_its_copy_lands( + kit, real_manager, attach_breaking({"_report_settled": settled_early}) + ) + + +@pytest.mark.parametrize( + "rows,usable", + [((0, 1), 64), ((0, 2), 32), ((), 0), ((1, 2), 0)], + ids=["prefix", "gap", "nothing", "no_first_block"], +) +def test_a_partial_arrival_is_usable_up_to_its_first_gap(kit, real_manager, rows, usable): + with real_manager() as mgr_a, real_manager() as mgr_b: + source = kit.published(mgr_a, SOURCE, PROMPT) + lender_a = attach(mgr_a, fetch_tokens=END) + publish = lender_a.lend_read(source, 0, END) + publish_view = ready_view(publish, mgr_a) + target = kit.admitted(mgr_b, TARGET, PROMPT) + lender_b = attach(mgr_b, fetch_tokens=END) + lease = lender_b.lend_write(target, 0, END) + view = lease.poll() + kit.fill_sentinel(mgr_b, target) + kit.relay(lender_a, publish_view, lender_b, view, deliver=lambda run, row: row in rows) + lease.mark_arrived((np.isin(np.arange(BLOCKS), rows),)) + mgr_b._stream.synchronize() + assert lender_b.readiness(target) == (usable, 0) + sentinel = bytes([kit.SENTINEL]) * kit.DevicePages(mgr_b).page_bytes(0) + for ordinal in range(BLOCKS): + expected = kit.page(mgr_a, source, 0, ordinal) if ordinal in rows else sentinel + got = kit.digest(kit.page(mgr_b, target, 0, ordinal)) + assert got == kit.digest(expected), "an unmarked row was copied" + lease.release() + publish.release() + + +@pytest.mark.parametrize("partial", [False, True], ids=["complete", "partial"]) +def test_a_windowed_fetch_is_usable_only_from_its_floor(kit, real_manager, partial): + windows = [WINDOW, 256] + with real_manager(windows=windows) as mgr_a, real_manager(windows=windows) as mgr_b: + source = kit.published(mgr_a, SOURCE, WINDOWED_PROMPT) + lender_a = attach(mgr_a, fetch_tokens=WINDOWED_END) + publish = lender_a.lend_read(source, 0, WINDOWED_END) + publish_view = ready_view(publish, mgr_a) + target = kit.admitted(mgr_b, TARGET, WINDOWED_PROMPT) + lender_b = attach(mgr_b, fetch_tokens=WINDOWED_END) + lease = lender_b.lend_write(target, 0, WINDOWED_END) + assert kit.kv(mgr_b, target).history_length == WINDOWED_END, "the history moves to the end" + view = lease.poll() + assert by_group(view) == by_group(publish_view) + kit.fill_sentinel(mgr_b, target) + deliver = (lambda run, row: row < 2) if partial else None + masks = kit.relay(lender_a, publish_view, lender_b, view, deliver) + lease.mark_arrived(masks) + mgr_b._stream.synchronize() + readiness = lender_b.readiness(target) + assert readiness.restart_floor == WINDOWED_END, "the window released the early blocks" + if partial: + assert readiness.usable_until < readiness.restart_floor, "empty: compute from 0" + else: + assert readiness.usable_until == WINDOWED_END + for run in view.runs: + for ordinal in run.ordinals.tolist(): + lg = run.layer_group + got = kit.digest(kit.page(mgr_b, target, lg, ordinal)) + assert got == kit.digest(kit.page(mgr_a, source, lg, ordinal)) + lease.release() + publish.release() + + +def test_mark_arrived_takes_one_mask_per_run_once_on_a_write_seen_ready(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + target = kit.admitted(mgr, TARGET, OTHER_PROMPT) + lender = attach(mgr, fetch_tokens=END, max_fetches=2) + read = lender.lend_read(source, 0, END) + read_view = ready_view(read, mgr) + with pytest.raises(RuntimeError): + read.mark_arrived(read_view.row_masks()) + write = lender.lend_write(target, 0, END) + with pytest.raises(RuntimeError): + write.mark_arrived(()) # before poll() gave the view + view = write.poll() + bad = ((), (np.ones(BLOCKS - 1, bool),), (np.ones(BLOCKS, bool),) * 2) + for masks in bad: + with pytest.raises(ValueError): + write.mark_arrived(masks) + assert lender.readiness(target) is None, "a refused mask recorded something" + write.mark_arrived(view.row_masks()) + with pytest.raises(RuntimeError): + write.mark_arrived(view.row_masks()) + mgr._stream.synchronize() + kv = kit.kv(mgr, target) + assert lender.readiness(target) == (kv.num_committed_tokens, kv.history_length) + write.release() + read.release() + + +@pytest.mark.parametrize( + "mark_first", [True, False], ids=["mark_then_release", "release_then_mark"] +) +def test_either_order_of_marks_and_release_copies_the_same(kit, real_manager, mark_first): + with real_manager() as mgr: + first = kit.admitted(mgr, TARGET, PROMPT) + second = kit.admitted(mgr, OTHER, OTHER_PROMPT) + lender = attach(mgr, fetch_tokens=END) # the second waits for the first's slots + lease = lender.lend_write(first, 0, END) + view = lease.poll() + kit.fill_sentinel(mgr, first) + kit.stage(lender, view, 0x5A) + waiting = lender.lend_write(second, 0, END) + if not mark_first: + lease.release() + assert waiting.poll() is None, "slots returned before the marks" + with kit.held_stream(mgr._stream) as gate: + lease.mark_arrived(view.row_masks(True)) + lease.release() + assert waiting.poll() is None, "slots returned before the copy out of them ran" + gate.open() + waiting_view = ready_view(waiting, mgr) + staged = bytes([0x5A]) * kit.DevicePages(mgr).page_bytes(0) + got = kit.digest([kit.page(mgr, first, 0, o) for o in range(BLOCKS)]) + assert got == kit.digest([staged] * BLOCKS) + assert lender.readiness(first) == (END, 0) + waiting.mark_arrived(waiting_view.row_masks()) + waiting.release() + + +# -- the request exits during a lease --------------------------------------------------------- + + +def check_arrivals_reach_only_pages_still_lent(kit, real_manager, attach): + with real_manager(max_tokens=kit.POOL_TOKENS) as mgr: + target = kit.admitted(mgr, TARGET, PROMPT) + lender = attach(mgr, fetch_tokens=END) + lease = lender.lend_write(target, 0, END) + view = lease.poll() + assert view is not None + kv = kit.kv(mgr, target) + lent = {kit.pages(kv, 0)[o] for o in view.runs[0].ordinals.tolist()} + kit.stage(lender, view, 0x5A) + mgr.free_resources(target) + with pytest.raises(ValueError): + lender.readiness(target) + others = kit.Requests(mgr) + try: + assert others.allocate(kit.pool_pages(mgr)) + assert lent <= others.pages(), "no lent page was reused: the check proves nothing" + lease.mark_arrived(view.row_masks(True)) + mgr._stream.synchronize() + dev = kit.DevicePages(mgr) + sentinel = bytes([kit.SENTINEL]) * dev.page_bytes(0) + for slot in sorted(lent): + got = kit.digest(dev.read(0, slot)) + assert got == kit.digest(sentinel), "an arrival overwrote another request's page" + finally: + lease.release() + others.free() + + +def test_arrivals_for_a_freed_request_copy_nothing(kit, real_manager): + check_arrivals_reach_only_pages_still_lent(kit, real_manager, attach) + + +def test_the_check_catches_a_lender_copying_into_pages_no_longer_lent(kit, real_manager): + def always_lent(self, kv, lease): + return [np.ones_like(m, dtype=bool) for m in real("_still_lent")(self, kv, lease)] + + with pytest.raises(CAUGHT, match="an arrival overwrote another request's page"): + check_arrivals_reach_only_pages_still_lent( + kit, real_manager, attach_breaking({"_still_lent": always_lent}) + ) + + +def check_a_free_fails_the_request_s_waiting_leases(kit, real_manager, attach): + with real_manager() as mgr: + holder = kit.published(mgr, SOURCE, PROMPT) + leaving = kit.published(mgr, OTHER, OTHER_PROMPT) + lender = attach(mgr, fetch_tokens=END) + head = lender.lend_read(holder, 0, END) + ready_view(head, mgr) + waiting = lender.lend_read(leaving, 0, 2 * TPB) + assert waiting.poll() is None and waiting.failure is None + mgr.free_resources(leaving) + assert waiting.failure is not None, "a lease waiting for slots outlived its request" + assert waiting.poll() is None + waiting.release() + head.release() + again = lender.lend_read(holder, 0, END) + ready_view(again, mgr) + again.release() + + +def test_a_free_fails_the_request_s_waiting_leases(kit, real_manager): + check_a_free_fails_the_request_s_waiting_leases(kit, real_manager, attach) + + +def test_the_check_catches_a_lender_keeping_a_freed_request_s_leases_in_line(kit, real_manager): + liar = attach_breaking({"_fail_waiting": lambda self, request_id, reason: None}) + with pytest.raises(CAUGHT, match="a lease waiting for slots outlived its request"): + check_a_free_fails_the_request_s_waiting_leases(kit, real_manager, liar) + + +def test_a_publish_granted_before_its_request_is_freed_copies_what_it_lent(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + original = [kit.page(mgr, source, 0, o) for o in range(BLOCKS)] + lender = attach(mgr, fetch_tokens=END) + with kit.held_stream(mgr._stream, strict=False) as gate: + lease = lender.lend_read(source, 0, END) # granted, its copy queued behind the gate + mgr.free_resources(source) + assert lease.failure is None and kit.kv(mgr, source) is None + gate.open() + (run,) = ready_view(lease, mgr).runs + length = lender.parts[run.part].slot_bytes + got = kit.digest([kit.host_bytes(a, length) for a in run.addresses.tolist()]) + assert got == kit.digest(original) + lease.release() + + +def test_a_waiting_read_fails_when_its_cache_changed_before_its_grant(kit, real_manager): + with real_manager() as mgr: + holder = kit.published(mgr, SOURCE, PROMPT) + other = kit.published(mgr, OTHER, OTHER_PROMPT) + lender = attach(mgr, fetch_tokens=END) + head = lender.lend_read(holder, 0, END) + ready_view(head, mgr) + waiting = lender.lend_read(other, 0, 2 * TPB) + assert waiting.poll() is None and waiting.failure is None + mgr.suspend_request(other) + head.release() + assert waiting.poll() is None and waiting.failure is not None, "copied a moved cache" + waiting.release() + again = lender.lend_read(holder, 0, END) # its slots came back + ready_view(again, mgr) + again.release() + + +# -- abandoned fetches ------------------------------------------------------------------------ + + +def check_an_abandoned_fetch_leaves_only_its_growth(kit, real_manager, attach, windowed): + tokens, other, end = ( + (WINDOWED_PROMPT, OTHER_WINDOWED_PROMPT, WINDOWED_END) + if windowed + else (PROMPT, OTHER_PROMPT, END) + ) + with real_manager(windows=[WINDOW, 256] if windowed else None) as mgr: + target = kit.admitted(mgr, TARGET, tokens) + second = kit.admitted(mgr, OTHER, other) + lender = attach(mgr, fetch_tokens=end) + lease = lender.lend_write(target, 0, end) + assert lender.readiness(target) is None + lease.release() # before any poll: no backend wrote, the fetch is abandoned + kv = kit.kv(mgr, target) + if windowed: + assert kv.history_length == end, "the growth moved the window's history" + own = (kv.num_committed_tokens, kv.history_length) + assert lender.readiness(target) == own, "the abandoned fetch still counts" + again = lender.lend_write(second, 0, end) + view = again.poll() + assert view is not None, "the abandoned fetch kept its slots" + again.mark_arrived(view.row_masks()) + again.release() + + +@pytest.mark.parametrize("windowed", [False, True], ids=["full", "windowed"]) +def test_an_abandoned_fetch_leaves_only_its_growth(kit, real_manager, windowed): + check_an_abandoned_fetch_leaves_only_its_growth(kit, real_manager, attach, windowed) + + +def test_the_check_catches_a_lender_keeping_an_abandoned_fetch(kit, real_manager): + liar = attach_breaking({"_abandon": lambda self, lease: None}) + with pytest.raises(CAUGHT, match="the abandoned fetch still counts"): + check_an_abandoned_fetch_leaves_only_its_growth(kit, real_manager, liar, windowed=True) + + +# -- a fetch split into leases, and a shrink of the cache it went into ------------------------- + +SPLIT_PROMPT = list(range(3000, 3129)) # four whole blocks and one token +SMALL_POOL = dict(max_tokens=256, max_batch_size=64) + + +def stale_anywhere(kit, mgr, history): + """Some windowed layer group has released blocks behind its window at ``history``.""" + for lg, window in enumerate(kit.windows(mgr)): + if window is not None: + beg, end = kit.stale_blocks(mgr, lg, history) + if end > beg: + return True + return False + + +def fetch_segment(kit, mgr, lender, target, start, end, byte, arrived=None): + """One write lease over ``[start, end)``: every slot staged with ``byte``, the rows of block + ordinals ``arrived`` (default all) marked, released, the copy into pages run.""" + lease = lender.lend_write(target, start, end) + view = lease.poll() + assert view is not None, f"the write lease was not ready: {lease.failure}" + kit.stage(lender, view, byte) + masks = tuple( + np.ones(len(run), dtype=bool) if arrived is None else np.isin(run.ordinals, arrived) + for run in view.runs + ) + lease.mark_arrived(masks) + lease.release() + mgr._stream.synchronize() + + +SEGMENTS = { + # name: (manager kwargs, segments [(start, end, arrived)], usable_until after each segment) + "two_halves": ({}, [(0, 64, None), (64, 128, None)], [64, 128]), + "three_segments": ({}, [(0, 32, None), (32, 96, None), (96, 128, None)], [32, 96, 128]), + "windowed": ({"windows": [WINDOW, 256]}, [(0, 64, None), (64, 128, None)], [64, 128]), + # A gap stays a gap across segments. + "first_landed_nothing": ({}, [(0, 64, []), (64, 128, None)], [0, 0]), + "first_landed_one_block": ({}, [(0, 64, [0]), (64, 128, None)], [32, 32]), +} + + +def check_segments_add_up(kit, real_manager, attach, case): + """Consecutive leases into one cache, nothing committed in between: readiness covers every + segment that landed, and nothing more.""" + manager_kwargs, segments, expected = SEGMENTS[case] + with real_manager(**manager_kwargs) as mgr: + target = kit.admitted(mgr, TARGET, SPLIT_PROMPT) + kv = kit.kv(mgr, target) + assert kv.num_committed_tokens == 0, "no local prefix: every usable token was fetched" + lender = attach(mgr, fetch_tokens=64) + bytes_of = {} + for i, ((start, end, arrived), usable) in enumerate(zip(segments, expected)): + byte = 0x11 * (i + 1) + fetch_segment(kit, mgr, lender, target, start, end, byte, arrived) + for ordinal in range(start // TPB, end // TPB): + if arrived is None or ordinal in arrived: + bytes_of[ordinal] = byte + assert kv.num_committed_tokens == 0, "a fetch committed" + readiness = lender.readiness(target) + floor = kv.history_length if stale_anywhere(kit, mgr, kv.history_length) else 0 + assert readiness is not None and readiness.restart_floor == floor + assert readiness.usable_until >= usable, ( + f"after segment {i} [{start}, {end}): readiness {tuple(readiness)}; an earlier " + "segment's delivered prefix was dropped" + ) + assert readiness.usable_until == usable, ( + f"after segment {i}: readiness {tuple(readiness)} counts what never landed" + ) + # The bytes are in the pages: only the bookkeeping decides whether they count. + dev = kit.DevicePages(mgr) + for lg, window in enumerate(kit.windows(mgr)): + if window is not None and window < 128: + continue + slots = kit.pages(kv, lg) + for ordinal, byte in bytes_of.items(): + got = kit.digest(dev.read(lg, slots[ordinal])) + assert got == kit.digest(bytes([byte]) * dev.page_bytes(lg)) + + +@pytest.mark.parametrize("case", list(SEGMENTS)) +def test_consecutive_leases_add_up_to_one_fetch(kit, real_manager, case): + check_segments_add_up(kit, real_manager, attach, case) + + +def test_the_check_catches_a_lender_keeping_only_the_latest_segment(kit, real_manager): + def latest_only(self, request_id, fetch, rows, copied): + self._delivered.pop(request_id, None) + real("_deliver")(self, request_id, fetch, rows, copied) + + with pytest.raises(CAUGHT, match="an earlier segment's delivered prefix was dropped"): + check_segments_add_up( + kit, real_manager, attach_breaking({"_deliver": latest_only}), "two_halves" + ) + + +@pytest.mark.parametrize("case", ["first_landed_nothing", "first_landed_one_block"]) +def test_the_check_catches_a_lender_counting_rows_that_never_landed(kit, real_manager, case): + def whole_ranges(self, manager, delivered, committed): # up to the last row any lease covered + return max(committed, max(len(b) for b in delivered.blocks) * TPB) + + with pytest.raises(CAUGHT, match="counts what never landed"): + check_segments_add_up( + kit, real_manager, attach_breaking({"_usable_until": whole_ranges}), case + ) + + +def check_an_abandoned_segment_keeps_earlier_deliveries(kit, real_manager, attach): + with real_manager() as mgr: + target = kit.admitted(mgr, TARGET, SPLIT_PROMPT) + lender = attach(mgr, fetch_tokens=64) + fetch_segment(kit, mgr, lender, target, 0, 64, 0x11) + assert lender.readiness(target) == (64, 0) + second = lender.lend_write(target, 64, 128) + assert lender.readiness(target) is None, "a fetch counts from its call" + second.release() # before anyone saw it ready: abandoned, nothing written + readiness = lender.readiness(target) + assert readiness == (64, 0), ( + f"readiness {tuple(readiness)} after an abandoned segment; the first segment's " + "delivered prefix was dropped" + ) + + +def test_an_abandoned_segment_keeps_what_earlier_ones_delivered(kit, real_manager): + check_an_abandoned_segment_keeps_earlier_deliveries(kit, real_manager, attach) + + +def test_the_check_catches_a_lender_dropping_every_delivery_on_abandon(kit, real_manager): + def abandon_all(self, lease): + real("_abandon")(self, lease) + self._delivered.pop(lease._request_id, None) + + with pytest.raises(CAUGHT, match="the first segment's delivered prefix was dropped"): + check_an_abandoned_segment_keeps_earlier_deliveries( + kit, real_manager, attach_breaking({"_abandon": abandon_all}) + ) + + +FRESH_FILL = 2.5 # every fp16 element 0x4100: no staged byte pattern reads as it + + +def test_a_resume_with_the_fresh_page_fill_on_keeps_the_fetched_blocks( + kit, real_manager, monkeypatch +): + """The manager's fresh-page fill writes pages as a request first gets them. The pages a fetch + grew are the request's own once the fetch lands: resuming fills only pages past them.""" + monkeypatch.setenv("TRTLLM_KV_FRESH_PAGE_FILL", str(FRESH_FILL)) + with real_manager() as mgr: + target = kit.admitted(mgr, TARGET, PROMPT) + lender = attach(mgr, fetch_tokens=END) + fetch_segment(kit, mgr, lender, target, 0, 64, 0x5A) + assert lender.readiness(target).usable_until == 64 + # The executor resumes the request where the fetch left it usable. + target.context_chunk_size = target.prompt_len - target.context_current_position + target.set_prepopulated_prompt_len(64, TPB) + target.context_chunk_size = target.prompt_len - 64 + assert mgr.resize_context(target, target.context_remaining_length) + mgr._stream.synchronize() + dev = kit.DevicePages(mgr) + fetched = kit.digest(bytes([0x5A]) * dev.page_bytes(0)) + filled = np.full(dev.page_bytes(0) // 2, FRESH_FILL, dtype=np.float16).tobytes() + got = [kit.digest(kit.page(mgr, target, 0, ordinal)) for ordinal in range(3)] + assert got[2] == kit.digest(filled), "the fill did not run on the resume" + assert got[:2] == [fetched, fetched], "the resume filled the fetched blocks as fresh pages" + mgr.free_resources(target) + + +def check_a_rollback_voids_the_delivery(kit, real_manager, attach, ask_between): + """The manager's own context rollback undoes the growth ``lend_write`` made: the same cache + shrinks and frees the fetched pages, which others overwrite before the cache grows back. + ``ask_between``: readiness is asked between the rollback and the regrow.""" + with real_manager() as mgr_a, real_manager(**SMALL_POOL) as mgr_b: + source = kit.published(mgr_a, SOURCE, PROMPT) + lender_a = attach(mgr_a, fetch_tokens=END) + publish = lender_a.lend_read(source, 0, END) + publish_view = ready_view(publish, mgr_a) + source_pages = [kit.digest(kit.page(mgr_a, source, 0, b)) for b in range(BLOCKS)] + target = kit.admitted(mgr_b, TARGET, PROMPT) + lender_b = attach(mgr_b, fetch_tokens=END) + kv = kit.kv(mgr_b, target) + lease = lender_b.lend_write(target, 0, END) + assert target.py_ctx_pre_resize_cap == 0, "the rollback would not undo the growth" + view = lease.poll() + kit.fill_sentinel(mgr_b, target) + lease.mark_arrived(kit.relay(lender_a, publish_view, lender_b, view)) + lease.release() + publish.release() + mgr_b._stream.synchronize() + assert [kit.digest(kit.page(mgr_b, target, 0, b)) for b in range(BLOCKS)] == source_pages + assert lender_b.readiness(target) == (END, 0) + + assert mgr_b.revert_allocate_context(target) is True + assert kit.kv(mgr_b, target) is kv, "the rollback replaced the cache" + assert kv.capacity == 0 and kit.pages(kv, 0) == [] and kv.num_committed_tokens == 0 + if ask_between: + readiness = lender_b.readiness(target) + assert readiness == (0, 0), f"readiness {tuple(readiness)} counts freed blocks" + assert kit.overwrite_free_pages(mgr_b) >= BLOCKS + assert mgr_b.resize_context(target, END) + regrown = [kit.digest(kit.page(mgr_b, target, 0, b)) for b in range(BLOCKS)] + assert regrown != source_pages, "the regrown pages still hold the fetched bytes" + readiness = lender_b.readiness(target) + assert readiness == (0, 0), ( + f"readiness {tuple(readiness)} counts blocks whose pages now hold {regrown}" + ) + mgr_b.free_resources(target) + + +@pytest.mark.parametrize("ask_between", [False, True], ids=["regrown_unasked", "asked_between"]) +def test_a_rollback_voids_what_the_fetch_delivered(kit, real_manager, ask_between): + check_a_rollback_voids_the_delivery(kit, real_manager, attach, ask_between) + + +def test_the_check_catches_a_lender_ignoring_the_manager_s_shrink(kit, real_manager): + liar = attach_breaking({"_on_shrink": lambda self, request_id, kv_cache: None}) + with pytest.raises(CAUGHT, match="counts blocks whose pages now hold"): + check_a_rollback_voids_the_delivery(kit, real_manager, liar, ask_between=False) + + +def test_the_check_catches_a_lender_never_voiding_delivered_rows(kit, real_manager): + liar = attach_breaking({"_void_past": lambda self, delivered, kv: None}) + with pytest.raises(CAUGHT, match="counts freed blocks"): + check_a_rollback_voids_the_delivery(kit, real_manager, liar, ask_between=True) + + +def check_a_rollback_to_a_local_prefix_keeps_only_the_prefix(kit, real_manager, attach): + """The rollback's other branch: a local match of two blocks, a fetch of the third; the rollback + shrinks to the prefix and suspends the cache.""" + local = 2 * TPB + with real_manager() as mgr_a, real_manager(**SMALL_POOL) as mgr_b: + source = kit.published(mgr_a, SOURCE, PROMPT) + lender_a = attach(mgr_a, fetch_tokens=END) + publish = lender_a.lend_read(source, 0, END) + publish_view = ready_view(publish, mgr_a) + kit.published(mgr_b, 7, PROMPT[:local] + list(range(9000, 9033))) + target = kit.admitted(mgr_b, TARGET, PROMPT) + lender_b = attach(mgr_b, fetch_tokens=END) + kv = kit.kv(mgr_b, target) + assert kv.num_committed_tokens == local + lease = lender_b.lend_write(target, local, END) + view = lease.poll() + kit.fill_sentinel(mgr_b, target, first_block=2) + lease.mark_arrived(kit.relay(lender_a, publish_view, lender_b, view)) + lease.release() + publish.release() + mgr_b._stream.synchronize() + assert lender_b.readiness(target) == (END, local) + assert mgr_b.revert_allocate_context(target) is True + assert kit.kv(mgr_b, target) is kv + assert (kv.capacity, kv.is_active, kv.num_committed_tokens) == (local, False, local) + readiness = lender_b.readiness(target) + assert readiness == (local, local), ( + f"readiness {tuple(readiness)} counts block 2, which the rollback freed" + ) + mgr_b.free_resources(target) + + +def test_a_rollback_to_a_local_prefix_keeps_only_the_prefix(kit, real_manager): + check_a_rollback_to_a_local_prefix_keeps_only_the_prefix(kit, real_manager, attach) + + +def test_the_check_catches_a_lender_counting_rows_past_the_shrunk_cache(kit, real_manager): + liar = attach_breaking( + { + "_on_shrink": lambda self, request_id, kv_cache: None, + "_void_past": lambda self, delivered, kv: None, + } + ) + with pytest.raises(CAUGHT, match="counts block 2, which the rollback freed"): + check_a_rollback_to_a_local_prefix_keeps_only_the_prefix(kit, real_manager, liar) + + +# -- failures at the call --------------------------------------------------------------------- + + +@pytest.mark.parametrize("state", ["no_cache", "suspended", "shut_down"]) +def test_a_lease_failed_at_the_call_changed_nothing(kit, real_manager, state): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + target = kit.admitted(mgr, TARGET, OTHER_PROMPT) + lender = attach(mgr, fetch_tokens=END, max_fetches=2) + if state == "no_cache": + source, target = kit.make_request(8, PROMPT), kit.make_request(9, OTHER_PROMPT) + elif state == "suspended": + for request in (source, target): + mgr.suspend_request(request) + assert not kit.kv(mgr, request).is_active + else: + mgr.shutdown() + before = [state_of(kit, mgr, lender, r) for r in (source, target)] + leases = [lender.lend_read(source, 0, END), lender.lend_write(target, 0, END)] + for lease in leases: + assert lease.failure is not None, "not failed at the call" + assert lease.poll() is None + assert [state_of(kit, mgr, lender, r) for r in (source, target)] == before + for lease in leases: + lease.release() + lease.release() + + +def test_a_write_the_pool_cannot_grow_fails_at_the_call_and_changed_nothing(kit, real_manager): + with real_manager(max_tokens=kit.POOL_TOKENS) as mgr: + target = kit.admitted(mgr, TARGET, PROMPT) + lender = attach(mgr, fetch_tokens=END) + others = kit.Requests(mgr) + others.allocate(kit.pool_pages(mgr)) + before = state_of(kit, mgr, lender, target) + lease = lender.lend_write(target, 0, END) + assert lease.failure is not None and lease.poll() is None + assert state_of(kit, mgr, lender, target) == before + lease.release() + others.free() + lease = lender.lend_write(target, 0, END) # with pages free again it goes through + view = lease.poll() + assert view is not None + lease.mark_arrived(view.row_masks()) + lease.release() + + +def test_a_write_left_with_a_block_without_a_page_fails_at_its_first_poll( + kit, real_manager, monkeypatch +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _manager + + with real_manager() as mgr: + target = kit.admitted(mgr, TARGET, PROMPT) + lender = attach(mgr, fetch_tokens=END) + kv = kit.kv(mgr, target) + + def without_block_one(read): + def pages(kv_cache, lg): + found = np.array(read(kv_cache, lg), dtype=np.int64) + if kv_cache is kv and len(found) > 1: + found[1] = -1 + return found + + return pages + + with monkeypatch.context() as patched: + for name in ("pages", "locked_pages"): + patched.setattr(_manager, name, without_block_one(getattr(_manager, name))) + lease = lender.lend_write(target, 0, END) + assert lease.failure is None, "the cache grew: nothing fails at the call after that" + assert lender.readiness(target) is None + assert lease.poll() is None and lease.failure is not None + assert lender.readiness(target) == (kv.num_committed_tokens, kv.history_length) + lease.release() + + +# -- argument errors -------------------------------------------------------------------------- + + +def test_a_read_refuses_a_bad_range_at_the_call(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=4 * TPB) + for start, end in ((-TPB, TPB), (TPB, 0), (0, 50), (16, 48), (0, 4 * TPB)): + with pytest.raises(ValueError): + lender.lend_read(source, start, end) + lease = lender.lend_read(source, 0, END) + ready_view(lease, mgr) + lease.release() + + +def test_a_write_refuses_a_bad_range_before_changing_anything(kit, real_manager): + with real_manager() as mgr: + kit.published(mgr, 7, PROMPT[: 2 * TPB] + list(range(9000, 9033))) + matched = kit.admitted(mgr, TARGET, PROMPT) # its first two blocks match locally + fresh = kit.admitted(mgr, OTHER, OTHER_PROMPT) + assert kit.kv(mgr, matched).num_committed_tokens == 2 * TPB + lender = attach(mgr, fetch_tokens=4 * TPB, max_fetches=2) + before = [state_of(kit, mgr, lender, r) for r in (matched, fresh)] + bad = [ + (fresh, -TPB, TPB), + (fresh, TPB, 0), + (fresh, 0, 50), + (fresh, 16, 48), + (matched, TPB, END), # starts inside the committed blocks + (fresh, 0, 4 * TPB), # past the whole blocks of the prompt + ] + for request, start, end in bad: + with pytest.raises(ValueError): + lender.lend_write(request, start, end) + assert [state_of(kit, mgr, lender, r) for r in (matched, fresh)] == before + leases = [lender.lend_write(matched, 2 * TPB, END), lender.lend_write(fresh, 0, END)] + with pytest.raises(ValueError): + lender.lend_write(fresh, 0, END) # one unsettled fetch per cache + for lease in leases: + view = lease.poll() + lease.mark_arrived(view.row_masks()) + lease.release() + + +def test_a_replaced_cache_takes_a_new_fetch(kit, real_manager): + with real_manager() as mgr: + target = kit.admitted(mgr, TARGET, PROMPT) + lender = attach(mgr, fetch_tokens=END, max_fetches=2) + first = lender.lend_write(target, 0, END) + first_view = first.poll() # unsettled: seen ready, not marked + mgr.free_resources(target) + target = kit.admitted(mgr, TARGET, PROMPT) + second = lender.lend_write(target, 0, END) + view = second.poll() + assert view is not None and lender.readiness(target) is None + second.mark_arrived(view.row_masks()) + first.mark_arrived(first_view.row_masks()) + for lease in (first, second): + lease.release() + + +def test_a_windowed_write_may_not_end_below_the_history(kit, real_manager): + with real_manager(windows=[WINDOW, 256]) as mgr: + target = kit.admitted(mgr, TARGET, WINDOWED_PROMPT) + lender = attach(mgr, fetch_tokens=WINDOWED_END) + lender.lend_write(target, 0, WINDOWED_END).release() # abandoned; the history stays + assert kit.kv(mgr, target).history_length == WINDOWED_END + before = state_of(kit, mgr, lender, target) + with pytest.raises(ValueError): + lender.lend_write(target, 0, END) + assert state_of(kit, mgr, lender, target) == before + + +def check_a_windowed_target_with_scratch_reuse_fails_at_the_call(kit, real_manager, attach): + with real_manager(windows=[WINDOW, 256], swa_scratch_reuse=True) as mgr: + target = kit.make_request(TARGET, WINDOWED_PROMPT) + assert mgr.prepare_context(target) + kv = kit.kv(mgr, target) + assert kv.enable_swa_scratch_reuse, "the target must start with scratch reuse on" + lender = attach(mgr, fetch_tokens=WINDOWED_END) + before = state_of(kit, mgr, lender, target) + try: + lease = lender.lend_write(target, 0, WINDOWED_END) + except RuntimeError as error: # the manager refusing the grow a lender let through + raise AssertionError(f"not failed at the call: {error}") from error + assert lease.failure is not None and "scratch" in lease.failure, "not failed at the call" + assert lease.poll() is None + assert state_of(kit, mgr, lender, target) == before + lease.release() + kv.enable_swa_scratch_reuse = False + lease = lender.lend_write(target, 0, WINDOWED_END) # with it off the write goes through + view = lease.poll() + assert view is not None + lease.mark_arrived(view.row_masks()) + lease.release() + + +def test_a_windowed_target_with_scratch_reuse_fails_at_the_call(kit, real_manager): + check_a_windowed_target_with_scratch_reuse_fails_at_the_call(kit, real_manager, attach) + + +def test_the_check_catches_a_lender_ignoring_scratch_reuse(kit, real_manager): + liar = attach_breaking({"_scratch_reuse_on": lambda self, kv: False}) + with pytest.raises(CAUGHT, match="not failed at the call"): + check_a_windowed_target_with_scratch_reuse_fails_at_the_call(kit, real_manager, liar) + + +class _TooBig(Exception): + pass + + +def _oversize_fails_the_lease(): + """Rules turning the refusal of a range staging cannot hold into a lease failed at the call.""" + nobody = SimpleNamespace(py_request_id=10**9) + + def check_fits(self, counts): + try: + real("_check_fits")(self, counts) + except ValueError as error: + raise _TooBig() from error + + def lend(name): + def call(self, request, start, end): + try: + return real(name)(self, request, start, end) + except _TooBig: + return real(name)(self, nobody, 0, 0) + + return call + + return { + "_check_fits": check_fits, + "lend_read": lend("lend_read"), + "lend_write": lend("lend_write"), + } + + +def check_a_range_staging_cannot_hold_is_refused(kit, real_manager, attach): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + target = kit.admitted(mgr, TARGET, OTHER_PROMPT) + lender = attach(mgr, fetch_tokens=2 * TPB) # two slots + before = state_of(kit, mgr, lender, target) + refusal(lambda: lender.lend_read(source, 0, END)) + refusal(lambda: lender.lend_write(target, 0, END)) + assert state_of(kit, mgr, lender, target) == before + # A range of fetch_tokens tokens fits exactly. + read = lender.lend_read(source, 0, 2 * TPB) + ready_view(read, mgr) + read.release() + write = lender.lend_write(target, 0, 2 * TPB) + view = write.poll() + assert view is not None + write.mark_arrived(view.row_masks()) + write.release() + + +def test_a_range_longer_than_staging_holds_is_refused(kit, real_manager): + check_a_range_staging_cannot_hold_is_refused(kit, real_manager, attach) + + +def test_the_check_catches_a_lender_failing_the_lease_instead(kit, real_manager): + liar = attach_breaking(_oversize_fails_the_lease()) + with pytest.raises(CAUGHT, match=NOT_REFUSED): + check_a_range_staging_cannot_hold_is_refused(kit, real_manager, liar) + + +# -- readiness -------------------------------------------------------------------------------- + + +def test_readiness_without_a_fetch_is_the_cache_s_own_interval(kit, real_manager): + with real_manager(windows=[WINDOW, 256]) as mgr: + source = kit.published(mgr, SOURCE, WINDOWED_PROMPT) + lender = attach(mgr, fetch_tokens=WINDOWED_END) + kv = kit.kv(mgr, source) + readiness = lender.readiness(source) + assert readiness == (max(kv.num_committed_tokens, kv.history_length), kv.history_length) + assert all(type(v) is int for v in readiness) + with pytest.raises(ValueError): + lender.readiness(kit.make_request(9, PROMPT)) + + +def test_readiness_waits_on_nothing_and_settles_by_being_asked(kit, real_manager): + with real_manager() as mgr: + target = kit.admitted(mgr, TARGET, PROMPT) + lender = attach(mgr, fetch_tokens=END) + lease = lender.lend_write(target, 0, END) + view = lease.poll() + with kit.held_stream(mgr._stream) as gate: + lease.mark_arrived(view.row_masks(True)) + assert lender.readiness(target) is None + assert lender.readiness(target) is None + gate.open() + assert kit.wait_until(lambda: lender.readiness(target) is not None) + assert lender.readiness(target) == (END, 0) + lease.release() + + +# -- shutdown --------------------------------------------------------------------------------- + + +def test_the_manager_s_shutdown_waits_for_the_lender_s_copies(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=END) + parts = lender.parts + with kit.gated_stream(mgr._stream, open_after=1.0) as gate: + lender.lend_read(source, 0, END).release() # its copy waits behind the gate + mgr.shutdown() + assert gate.opened, "shutdown returned while a copy into staging was queued" + assert not kit.staging_kept(parts), "no lease was open: the memory is freed" + + +def test_shutdown_fails_waiting_leases_and_later_calls_only_end_records(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + target = kit.admitted(mgr, TARGET, OTHER_PROMPT) + lender = attach(mgr, fetch_tokens=END) + held = lender.lend_read(source, 0, END) + ready_view(held, mgr) + waiting = lender.lend_read(source, 0, END) + assert waiting.poll() is None + mgr.shutdown() + assert waiting.failure is not None and waiting.poll() is None + for lease in (lender.lend_read(source, 0, END), lender.lend_write(target, 0, END)): + assert lease.failure is not None and lease.poll() is None + lease.release() + with pytest.raises(ValueError): + lender.readiness(source) + for lease in (held, held, waiting): + lease.release() + mgr.shutdown() # a second shutdown frees nothing and returns + + +def check_an_open_lease_keeps_the_staging_memory(kit, real_manager, attach, monkeypatch, failed): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=END) + parts = lender.parts + assert kit.staging_memory(parts) + if failed: + lease = lender.lend_read(kit.make_request(9, PROMPT), 0, END) + assert lease.failure is not None + address, content = None, None + else: + lease = lender.lend_read(source, 0, END) + address = int(ready_view(lease, mgr).runs[0].addresses[0]) + content = kit.host_bytes(address, parts[0].slot_bytes) + warnings = kit.lender_warnings(monkeypatch) + mgr.shutdown() + assert kit.staging_kept(parts), "staging freed at shutdown with a lease open" + assert warnings, "keeping the memory until exit is logged" + if address is not None: + got = kit.digest(kit.host_bytes(address, parts[0].slot_bytes)) + assert got == kit.digest(content), "the kept staging memory changed" + lease.release() + assert kit.staging_kept(parts), "a release after the shutdown freed it" + + +@pytest.mark.parametrize("failed", [False, True], ids=["ready_lease", "failed_lease"]) +def test_an_unreleased_lease_keeps_the_staging_memory_until_exit( + kit, real_manager, monkeypatch, failed +): + check_an_open_lease_keeps_the_staging_memory(kit, real_manager, attach, monkeypatch, failed) + + +def test_the_check_catches_a_lender_freeing_memory_a_lease_holds(kit, real_manager, monkeypatch): + liar = attach_breaking({"_memory_in_use": lambda self: False}) + with pytest.raises(CAUGHT, match="staging freed at shutdown with a lease open"): + check_an_open_lease_keeps_the_staging_memory( + kit, real_manager, liar, monkeypatch, failed=False + ) + + +def test_with_every_lease_released_shutdown_frees_the_staging_memory( + kit, real_manager, monkeypatch +): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=END) + parts = lender.parts + refs = [weakref.ref(m) for m in kit.staging_memory(parts)] + assert refs + lease = lender.lend_read(source, 0, END) + ready_view(lease, mgr) + lease.release() + lender.lend_read(kit.make_request(9, PROMPT), 0, END).release() # failed, released + warnings = kit.lender_warnings(monkeypatch) + mgr.shutdown() + assert not kit.staging_kept(parts) and warnings == [] + gc.collect() + assert all(ref() is None for ref in refs), "something kept the staging memory alive" + + +def test_a_shutdown_hook_that_fails_keeps_the_memory_and_raises_nothing( + kit, real_manager, monkeypatch +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + def planted(self): + raise RuntimeError("planted") + + with real_manager() as mgr: + parts = attach(mgr, fetch_tokens=END).parts + monkeypatch.setattr(_lender.Staging, "_memory_in_use", planted) + mgr.shutdown() + assert kit.staging_kept(parts) + + +# -- parts holds ------------------------------------------------------------------------------ + + +def released(hold) -> None: + """``hold.release()``; the check fails if it raises, since release is legal in every state.""" + try: + hold.release() + except Exception as error: # handed to the check as its failure + raise AssertionError(f"release raised {error!r}") from error + + +def check_an_open_hold_keeps_the_staging_memory(kit, real_manager, attach, monkeypatch): + with real_manager() as mgr: + lender = attach(mgr, fetch_tokens=END) + parts = lender.parts + assert kit.staging_memory(parts) + hold = lender.hold_parts() + assert isinstance(hold, PartsHold) + warnings = kit.lender_warnings(monkeypatch) + mgr.shutdown() + assert kit.staging_kept(parts), "staging freed at shutdown with a hold open" + assert warnings, "keeping the memory until exit is logged" + released(hold) + mgr.shutdown() + assert kit.staging_kept(parts), "a release after the shutdown freed it" + + +def test_an_open_hold_keeps_the_staging_memory_until_exit(kit, real_manager, monkeypatch): + check_an_open_hold_keeps_the_staging_memory(kit, real_manager, attach, monkeypatch) + + +def test_the_check_catches_a_lender_ignoring_holds(kit, real_manager, monkeypatch): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + liar = attach_breaking({"hold_parts": lambda self: _lender._PartsHold(None)}) + with pytest.raises(CAUGHT, match="staging freed at shutdown with a hold open"): + check_an_open_hold_keeps_the_staging_memory(kit, real_manager, liar, monkeypatch) + + +def check_a_hold_released_before_shutdown_lets_the_memory_go( + kit, real_manager, attach, monkeypatch +): + with real_manager() as mgr: + lender = attach(mgr, fetch_tokens=END) + parts = lender.parts + refs = [weakref.ref(m) for m in kit.staging_memory(parts)] + assert refs + lender.hold_parts().release() + warnings = kit.lender_warnings(monkeypatch) + mgr.shutdown() + assert not kit.staging_kept(parts), "a released hold kept the memory" + assert warnings == [] + gc.collect() + assert all(ref() is None for ref in refs), "something kept the staging memory alive" + + +def test_a_hold_released_before_shutdown_lets_the_memory_go(kit, real_manager, monkeypatch): + check_a_hold_released_before_shutdown_lets_the_memory_go(kit, real_manager, attach, monkeypatch) + + +def test_the_check_catches_a_lender_never_letting_a_hold_go(kit, real_manager, monkeypatch): + liar = attach_breaking({"_end_hold": lambda self, hold: None}) + with pytest.raises(CAUGHT, match="a released hold kept the memory"): + check_a_hold_released_before_shutdown_lets_the_memory_go( + kit, real_manager, liar, monkeypatch + ) + + +def check_a_dropped_hold_keeps_the_staging_memory(kit, real_manager, attach): + with real_manager() as mgr: + lender = attach(mgr, fetch_tokens=END) + parts = lender.parts + assert kit.staging_memory(parts) + lender.hold_parts() # dropped unreleased + gc.collect() + mgr.shutdown() + assert kit.staging_kept(parts), "a dropped hold let the memory go" + + +def test_a_hold_dropped_unreleased_still_keeps_the_staging_memory(kit, real_manager): + check_a_dropped_hold_keeps_the_staging_memory(kit, real_manager, attach) + + +def test_the_check_catches_a_lender_holding_holds_weakly(kit, real_manager): + def weakly(self, *args): + real("__init__")(self, *args) + self._holds = weakref.WeakSet() + + with pytest.raises(CAUGHT, match="a dropped hold let the memory go"): + check_a_dropped_hold_keeps_the_staging_memory( + kit, real_manager, attach_breaking({"__init__": weakly}) + ) + + +def check_a_hold_releases_once_in_every_state(kit, real_manager, attach): + with real_manager() as mgr: + lender = attach(mgr, fetch_tokens=END) + parts = lender.parts + assert kit.staging_memory(parts) + first, second = lender.hold_parts(), lender.hold_parts() + released(first) + released(first) # does nothing: the second hold still holds + mgr.shutdown() + assert kit.staging_kept(parts), "a second release let the other hold go" + late = lender.hold_parts() # after the shutdown: inert + assert isinstance(late, PartsHold) + for hold in (late, late): + released(hold) + gone = weakref.ref(lender) + del mgr, lender + gc.collect() + assert gone() is None, "the lender outlived its manager" + for hold in (second, second): # its lender is gone + released(hold) + + +def test_a_hold_releases_once_in_every_state(kit, real_manager): + check_a_hold_releases_once_in_every_state(kit, real_manager, attach) + + +def test_the_check_catches_a_hold_releasing_again(kit, real_manager, monkeypatch): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + def unguarded(self): + lender = self._lender() if self._lender is not None else None + if lender is not None: + lender._end_hold(self) + + monkeypatch.setattr(_lender._PartsHold, "release", unguarded) + with pytest.raises(CAUGHT, match="release raised"): + check_a_hold_releases_once_in_every_state(kit, real_manager, attach) + + +class _CountedHold: + """A hold whose every release reaches its lender, which counts holds with an integer.""" + + def __init__(self, lender): + self._lender = weakref.ref(lender) + + def release(self): + lender = self._lender() + if lender is not None: + lender._end_hold(self) + + +def _counting_holds(self, *args): + real("__init__")(self, *args) + self._count = 0 + + +def _take_counted(self): + self._count += 1 + return _CountedHold(self) + + +def _drop_counted(self, hold): + self._count -= 1 + + +def _in_use_counted(self): + return bool(self._unreleased) or self._count > 0 or bool(self._quarantined) + + +def test_the_check_catches_a_lender_counting_holds_with_an_integer(kit, real_manager): + liar = attach_breaking( + { + "__init__": _counting_holds, + "hold_parts": _take_counted, + "_end_hold": _drop_counted, + "_memory_in_use": _in_use_counted, + } + ) + with pytest.raises(CAUGHT, match="a second release let the other hold go"): + check_a_hold_releases_once_in_every_state(kit, real_manager, liar) + + +# What is released before the manager's shutdown, in order; the rest is released after it. +ORDERS = { + "neither": (), + "lease": ("lease",), + "hold": ("hold",), + "lease_then_hold": ("lease", "hold"), + "hold_then_lease": ("hold", "lease"), +} + + +def check_a_lease_and_a_hold_keep_the_memory_until_both_go( + kit, real_manager, attach, monkeypatch, order +): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=END) + parts = lender.parts + assert kit.staging_memory(parts) + hold = lender.hold_parts() + lease = lender.lend_read(source, 0, END) + ready_view(lease, mgr) + ends = {"lease": lease.release, "hold": lambda: released(hold)} + for what in order: + ends[what]() + warnings = kit.lender_warnings(monkeypatch) + mgr.shutdown() + both = len(order) == 2 + if both: + assert not kit.staging_kept(parts) and warnings == [], "kept with both released" + else: + assert kit.staging_kept(parts), f"freed at shutdown after releasing {order}" + assert warnings, "keeping the memory until exit is logged" + for what in [w for w in ("lease", "hold") if w not in order]: + ends[what]() + assert kit.staging_kept(parts) == (not both), "a release after the shutdown changed it" + + +@pytest.mark.parametrize("order", list(ORDERS)) +def test_a_lease_and_a_hold_keep_the_memory_until_both_go(kit, real_manager, monkeypatch, order): + check_a_lease_and_a_hold_keep_the_memory_until_both_go( + kit, real_manager, attach, monkeypatch, ORDERS[order] + ) + + +def _in_every_order(check, *args): + for order in ORDERS.values(): + check(*args, order) + + +def test_the_check_catches_a_hold_release_that_forgets_open_leases(kit, real_manager, monkeypatch): + def end_hold(self, hold): + real("_end_hold")(self, hold) + self._unreleased.clear() + + liar = attach_breaking({"_end_hold": end_hold}) + with pytest.raises(CAUGHT, match="freed at shutdown"): + _in_every_order( + check_a_lease_and_a_hold_keep_the_memory_until_both_go, + kit, + real_manager, + liar, + monkeypatch, + ) + + +def test_the_check_catches_a_lease_release_that_forgets_open_holds(kit, real_manager, monkeypatch): + def on_release(self, lease): + real("_on_release")(self, lease) + self._holds.clear() + + def end_hold(self, hold): # a hold the lease's release already ended + self._holds.discard(hold) + + liar = attach_breaking({"_on_release": on_release, "_end_hold": end_hold}) + with pytest.raises(CAUGHT, match="freed at shutdown"): + _in_every_order( + check_a_lease_and_a_hold_keep_the_memory_until_both_go, + kit, + real_manager, + liar, + monkeypatch, + ) + + +def check_a_failed_copy_and_a_hold_keep_the_memory_in_every_order( + kit, real_manager, attach, monkeypatch, order +): + def no_event(self, stream=None): + raise RuntimeError("planted") + + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=END) + parts = lender.parts + assert kit.staging_memory(parts) + hold = lender.hold_parts() + with monkeypatch.context() as patched: + patched.setattr(torch.cuda.Event, "record", no_event) + lease = lender.lend_read(source, 0, END) # its copy has no event: slots lost + assert lease.failure is not None + mgr._stream.synchronize() + ends = {"lease": lease.release, "hold": lambda: released(hold)} + for what in order: + ends[what]() + warnings = kit.lender_warnings(monkeypatch) + mgr.shutdown() + assert kit.staging_kept(parts), f"freed at shutdown after releasing {order}" + assert warnings, "keeping the memory until exit is logged" + for what in [w for w in ("lease", "hold") if w not in order]: + ends[what]() + assert kit.staging_kept(parts), "a release after the shutdown freed it" + + +@pytest.mark.parametrize("order", list(ORDERS)) +def test_a_failed_copy_and_a_hold_keep_the_memory_in_every_order( + kit, real_manager, monkeypatch, order +): + check_a_failed_copy_and_a_hold_keep_the_memory_in_every_order( + kit, real_manager, attach, monkeypatch, ORDERS[order] + ) + + +def test_the_check_catches_a_hold_release_that_forgets_lost_slots(kit, real_manager, monkeypatch): + def end_hold(self, hold): + real("_end_hold")(self, hold) + self._quarantined.clear() + + liar = attach_breaking({"_end_hold": end_hold}) + with pytest.raises(CAUGHT, match="freed at shutdown"): + _in_every_order( + check_a_failed_copy_and_a_hold_keep_the_memory_in_every_order, + kit, + real_manager, + liar, + monkeypatch, + ) + + +def test_staging_outlives_a_manager_dropped_without_shutdown(kit): + torch.cuda.init() + gc.collect() + mgr = kit.make_manager() + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=END) + parts = lender.parts + lease = lender.lend_read(source, 0, END) + address = int(ready_view(lease, mgr).runs[0].addresses[0]) + content = kit.host_bytes(address, parts[0].slot_bytes) + mgr.free_resources(source) + mgr._stream.synchronize() + watched = weakref.ref(mgr) + del mgr, lender + gc.collect() + torch.cuda.empty_cache() + assert watched() is None, "an open lease kept the lender or the manager alive" + assert kit.staging_kept(parts) + got = kit.digest(kit.host_bytes(address, parts[0].slot_bytes)) + assert got == kit.digest(content), "the kept staging memory changed" + lease.release() + lease.release() + + +# -- progress --------------------------------------------------------------------------------- + + +class EventQueries: + """Counts ``torch.cuda.Event.query`` per event within each lender call made through ``call``.""" + + def __init__(self, monkeypatch): + self.asked: List[int] = [] + self.last: Set[int] = set() # the events the latest call asked + query = torch.cuda.Event.query + + def counted(event): + self.asked.append(id(event)) + return query(event) + + monkeypatch.setattr(torch.cuda.Event, "query", counted) + + def call(self, method, *args): + first = len(self.asked) + result = method(*args) + asked = self.asked[first:] + assert len(asked) == len(set(asked)), "a copy's event was asked twice in one lender call" + self.last = set(asked) + return result + + +HELD_READS = 4 + + +def check_each_call_asks_each_copy_once(kit, real_manager, attach, monkeypatch): + """Each lender call asks each copy's event at most once, a lease still lent costs no query + beyond its own poll's, and a copy reported complete is never asked again: a call asks only its + own copy and the pending copies of released leases, not one per held lease.""" + with real_manager() as mgr: + sources = [ + kit.published(mgr, 10 + i, [10000 * (i + 1) + t for t in range(END + 1)]) + for i in range(HELD_READS) + ] + target = kit.admitted(mgr, TARGET, PROMPT) + assert mgr._resize_for_connector_prefix(target, kit.kv(mgr, target), 0, END) + lender = attach(mgr, fetch_tokens=END, max_fetches=HELD_READS + 1) + events = EventQueries(monkeypatch) + with kit.held_stream(mgr._stream) as gate: + reads = [events.call(lender.lend_read, source, 0, END) for source in sources] + write = events.call(lender.lend_write, target, 0, END) + view = events.call(write.poll) + events.call(write.mark_arrived, view.row_masks(True)) + events.call(write.release) + released = id(write._copy._event) + for _ in range(2): + for lease in reads: + assert events.call(lease.poll) is None + own = id(lease._copy._event) + assert events.last == {own, released}, ( + f"a read's poll asked {len(events.last)} copies, not its own and the " + "released one" + ) + assert events.call(lender.readiness, target) is None + assert events.last == {released}, ( + f"readiness asked {len(events.last)} copies, not the released one" + ) + gate.open() + mgr._stream.synchronize() + assert events.call(lender.readiness, target) is not None + for lease in reads: + assert events.call(lease.poll) is not None + first = len(events.asked) + for _ in range(3): + for lease in reads: + assert events.call(lease.poll) is not None + assert events.call(lender.readiness, target) is not None + assert len(events.asked) == first, "a copy reported complete was asked again" + for lease in reads: + events.call(lease.release) + assert len(events.asked) == first, "a copy reported complete was asked again" + + +def test_each_call_asks_each_copy_s_event_at_most_once(kit, real_manager, monkeypatch): + check_each_call_asks_each_copy_once(kit, real_manager, attach, monkeypatch) + + +def test_the_check_catches_a_lender_asking_every_held_lease_on_every_call( + kit, real_manager, monkeypatch +): + def asks_first(self, lease): + if lease._copy is not None: + lease._copy._event.query() + return real("_recyclable")(self, lease) + + with pytest.raises(CAUGHT, match="asked twice in one lender call"): + check_each_call_asks_each_copy_once( + kit, real_manager, attach_breaking({"_recyclable": asks_first}), monkeypatch + ) + + +def test_the_check_catches_a_lender_asking_every_held_lease_s_copy_once_per_call( + kit, real_manager, monkeypatch +): + def lent_first(self, lease): # asks through the per-round cache, then applies the rule + if not self._landed(lease._copy): + return False + return real("_recyclable")(self, lease) + + with pytest.raises(CAUGHT, match="a read's poll asked 5 copies"): + check_each_call_asks_each_copy_once( + kit, real_manager, attach_breaking({"_recyclable": lent_first}), monkeypatch + ) + + +def test_the_check_catches_a_lender_asking_a_completed_copy_again(kit, real_manager, monkeypatch): + def forgets_completion(self, copy): # asks once per round, whatever it learnt before + if copy is not None and copy._round != self._round: + copy._round = self._round + copy._done = bool(copy._event.query()) + return copy is None or copy._done + + with pytest.raises(CAUGHT, match="a copy reported complete was asked again"): + check_each_call_asks_each_copy_once( + kit, real_manager, attach_breaking({"_landed": forgets_completion}), monkeypatch + ) + + +def check_progress_needs_no_new_traffic(kit, real_manager, attach): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=END) + holder = lender.lend_read(source, 0, END) + ready_view(holder, mgr) + waiting = lender.lend_read(source, 0, END) + assert waiting.poll() is None + holder.release() + ready_view(waiting, mgr) # its own polls alone grant it and see its copy land + waiting.release() + + +def test_progress_needs_no_new_traffic(kit, real_manager): + check_progress_needs_no_new_traffic(kit, real_manager, attach) + + +def test_the_check_catches_a_lender_progressing_only_when_lending(kit, real_manager): + lending = [0] + + def lend(name): + def call(self, *args): + lending[0] += 1 + try: + return real(name)(self, *args) + finally: + lending[0] -= 1 + + return call + + def progress(self): + if lending[0]: + real("_progress")(self) + + liar = attach_breaking( + {"_progress": progress, "lend_read": lend("lend_read"), "lend_write": lend("lend_write")} + ) + with pytest.raises(CAUGHT, match="the lease never became ready"): + check_progress_needs_no_new_traffic(kit, real_manager, liar) + + +# -- a copy that fails partway ---------------------------------------------------------------- + + +class FailingDriver: + """The CUDA driver, except that its ``fail_on``-th memcpy call, of any flavour, fails.""" + + def __init__(self, real_driver, fail_on): + self._real = real_driver + self._fail_on = fail_on + self.calls = 0 + + def __getattr__(self, name): + found = getattr(self._real, name) + if not name.startswith("cuMemcpy"): + return found + + def call(*args, **kwargs): + self.calls += 1 + if self.calls == self._fail_on: + return (self._real.CUresult.CUDA_ERROR_INVALID_VALUE,) + return found(*args, **kwargs) + + return call + + +def test_a_publish_whose_copy_fails_partway_fails_and_frees_its_slots_after_the_rest( + kit, real_manager, monkeypatch +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + # Three layers: two pool groups, so the copy is several memcpys. + with real_manager(windows=[WINDOW, 256], num_layers=3) as mgr: + source = kit.published(mgr, SOURCE, WINDOWED_PROMPT) + lender = attach(mgr, fetch_tokens=WINDOWED_END) + groups = kit.pool_group_ids(mgr) + driver = FailingDriver(_lender.drv, fail_on=2) + monkeypatch.setattr(_lender, "drv", driver) + with kit.held_stream(mgr._stream) as gate: + lease = lender.lend_read(source, 0, WINDOWED_END) # every slot + assert driver.calls >= 2, "a single memcpy: nothing failed partway" + assert lease.failure is not None and lease.poll() is None + lease.release() + assert [lender._free_slots(g) for g in groups] == [0] * len(groups), ( + "slots came back while a copy into them was queued" + ) + gate.open() + lender.readiness(source) # any call makes progress + assert [lender._free_slots(g) for g in groups] == [p.slots for p in lender.parts] + + +def test_marks_whose_copy_fails_partway_abandon_the_fetch_and_keep_the_lease_ready( + kit, real_manager, monkeypatch +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + with real_manager(windows=[WINDOW, 256], num_layers=3) as mgr: + target = kit.admitted(mgr, TARGET, WINDOWED_PROMPT) + lender = attach(mgr, fetch_tokens=WINDOWED_END) + groups = kit.pool_group_ids(mgr) + lease = lender.lend_write(target, 0, WINDOWED_END) + view = lease.poll() + kit.stage(lender, view, 0x5A) + driver = FailingDriver(_lender.drv, fail_on=2) + monkeypatch.setattr(_lender, "drv", driver) + kv = kit.kv(mgr, target) + with kit.held_stream(mgr._stream) as gate: + lease.mark_arrived(view.row_masks(True)) # raises nothing + assert driver.calls >= 2, "a single memcpy: nothing failed partway" + assert lender.readiness(target) == (kv.num_committed_tokens, kv.history_length) + assert lease.poll() is view, "the lease stays ready" + lease.release() + assert [lender._free_slots(g) for g in groups] == [0] * len(groups) + gate.open() + lender.readiness(target) + assert [lender._free_slots(g) for g in groups] == [p.slots for p in lender.parts] + + +def test_a_copy_without_an_event_loses_its_slots_and_keeps_the_memory( + kit, real_manager, monkeypatch +): + def no_event(self, stream=None): + raise RuntimeError("planted") + + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=END) + parts = lender.parts + (group,) = kit.pool_group_ids(mgr) + with monkeypatch.context() as patched: + patched.setattr(torch.cuda.Event, "record", no_event) + lease = lender.lend_read(source, 0, END) + assert lease.failure is not None + mgr._stream.synchronize() + lease.release() + lender.readiness(source) + assert lender._free_slots(group) == 0, "slots a copy may still touch were reused" + warnings = kit.lender_warnings(monkeypatch) + mgr.shutdown() + assert kit.staging_kept(parts) and warnings + + +# -- errors inside the lender ----------------------------------------------------------------- + + +def test_marks_that_raise_leave_the_fetch_abandoned_and_the_slots_returnable( + kit, real_manager, monkeypatch +): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _lender + + def planted(*args): + raise RuntimeError("planted") + + with real_manager() as mgr: + target = kit.admitted(mgr, TARGET, PROMPT) + other = kit.admitted(mgr, OTHER, OTHER_PROMPT) + lender = attach(mgr, fetch_tokens=END) + (group,) = kit.pool_group_ids(mgr) + lease = lender.lend_write(target, 0, END) + view = lease.poll() + kv = kit.kv(mgr, target) + with monkeypatch.context() as patched: + patched.setattr(_lender.Staging, "_still_lent", planted) + with pytest.raises(RuntimeError, match="planted"): + lease.mark_arrived(view.row_masks(True)) + assert lender.readiness(target) == (kv.num_committed_tokens, kv.history_length) + lease.release() + waiting = lender.lend_write(other, 0, END) + waiting_view = waiting.poll() + assert waiting_view is not None, "the lease's slots never came back" + opened = lender._open_count() + with monkeypatch.context() as patched: + patched.setattr(_lender.Staging, "_progress", planted) + try: + waiting.release() + except RuntimeError: + pass # the planted error may surface; the release must count all the same + assert lender._open_count() == opened - 1 + waiting.mark_arrived(waiting_view.row_masks()) + lender.readiness(other) + assert lender._free_slots(group) == lender.parts[0].slots + + +def test_a_grant_that_raises_fails_only_its_own_lease(kit, real_manager, monkeypatch): + from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import _manager + + with real_manager() as mgr: + holder = kit.published(mgr, SOURCE, PROMPT) + broken = kit.published(mgr, OTHER, OTHER_PROMPT) + fetched = kit.admitted(mgr, TARGET, list(range(8000, 8097))) + lender = attach(mgr, fetch_tokens=END) + write = lender.lend_write(fetched, 0, END) + write.mark_arrived(write.poll().row_masks()) + write.release() + head = lender.lend_read(holder, 0, END) + ready_view(head, mgr) + waiting = lender.lend_read(broken, 0, 2 * TPB) + assert waiting.poll() is None + broken_kv = kit.kv(mgr, broken) + + real_kv_of = _manager.kv_of + + def kv_of(manager, request_id): + if request_id == OTHER: + raise RuntimeError("planted") + return real_kv_of(manager, request_id) + + def raising_for_broken(read): + def call(kv_cache, *rest): + if kv_cache is broken_kv: + raise RuntimeError("planted") + return read(kv_cache, *rest) + + return call + + with monkeypatch.context() as patched: + patched.setattr(_manager, "kv_of", kv_of) + for name in ("cache_state", "pages", "locked_pages"): + patched.setattr(_manager, name, raising_for_broken(getattr(_manager, name))) + head.release() # grants the one waiting, whose grant raises + assert waiting.poll() is None and waiting.failure is not None + assert isinstance(lender.readiness(fetched), Readiness) + waiting.release() + again = lender.lend_read(holder, 0, END) + ready_view(again, mgr) + again.release() + + +# -- names ------------------------------------------------------------------------------------ + + +def _tp2(rank): + from tensorrt_llm.mapping import Mapping + + return Mapping(world_size=2, tp_size=2, rank=rank) + + +def names_by_rank(kit, real_manager, attach, num_kv_heads, scopes=(SCOPE, SCOPE)): + """Each rank of a TP2 pair publishes ``PROMPT``: its names, and its part names.""" + names, part_names = [], [] + for rank, scope in enumerate(scopes): + with real_manager(num_kv_heads=num_kv_heads, mapping=_tp2(rank)) as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=END, scope=scope) + lease = lender.lend_read(source, 0, END) + (run,) = ready_view(lease, mgr).runs + names.append(np.array(run.names)) + part_names.append([p.name for p in lender.parts]) + lease.release() + return names, part_names + + +def check_a_replicated_group_is_named_alike_on_every_rank(kit, real_manager, attach): + (rank0, rank1), (parts0, parts1) = names_by_rank(kit, real_manager, attach, num_kv_heads=1) + assert rank0.shape == (BLOCKS, 54) and rank0.dtype == np.uint8 + assert np.array_equal(rank0, rank1), "ranks holding the same bytes must share their names" + assert parts0 == parts1 + + +def test_a_replicated_group_is_named_alike_on_every_rank(kit, real_manager): + check_a_replicated_group_is_named_alike_on_every_rank(kit, real_manager, attach) + + +def test_the_check_catches_a_lender_naming_each_rank_apart(kit, real_manager): + def per_rank(mgr): + shard = np.frombuffer( + mgr.mapping.tp_size.to_bytes(2, "big") + mgr.mapping.tp_rank.to_bytes(2, "big"), + np.uint8, + ) + + def names(self, layer_group, keys): + out = np.array(real("_names")(self, layer_group, keys)) + out[:, 50:54] = shard + return out + + return {"_names": names} + + with pytest.raises(CAUGHT, match="ranks holding the same bytes must share their names"): + check_a_replicated_group_is_named_alike_on_every_rank( + kit, real_manager, attach_breaking(per_rank) + ) + + +def test_a_head_sharded_group_names_each_rank_s_share(kit, real_manager): + (rank0, rank1), (parts0, parts1) = names_by_rank(kit, real_manager, attach, num_kv_heads=4) + assert np.array_equal(rank0[:, :50], rank1[:, :50]), "same scope, layout and blocks" + for rank, names in enumerate((rank0, rank1)): + shard = (2).to_bytes(2, "big") + rank.to_bytes(2, "big") + assert {bytes(n[50:54]) for n in names} == {shard} + assert parts0 == parts1 + + +def test_a_different_scope_shares_no_name(kit, real_manager): + names = [] + for scope in (b"model-a", b"model-b"): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + lender = attach(mgr, fetch_tokens=END, scope=scope) + lease = lender.lend_read(source, 0, END) + names.append(([bytes(n) for n in ready_view(lease, mgr).runs[0].names], lender.parts)) + lease.release() + (first, first_parts), (second, second_parts) = names + assert [n[16:] for n in first] == [n[16:] for n in second], "same blocks, same keys" + assert not set(first) & set(second) + assert [p.name for p in first_parts] == [p.name for p in second_parts], "laid out alike" + + +def test_part_names_follow_the_layout_alone(real_manager): + names = [] + for tokens_per_block, scope in ((TPB, b"model-a"), (TPB, b"model-b"), (2 * TPB, b"model-a")): + with real_manager(tokens_per_block=tokens_per_block) as mgr: + lender = attach(mgr, fetch_tokens=END, scope=scope) + names.append([p.name for p in lender.parts]) + assert names[0] == names[1], "instances laid out alike name their parts alike" + assert names[0] != names[2], "another layout, other part names" + + +# -- threads ---------------------------------------------------------------------------------- + + +def test_threads_take_turns_and_the_lender_starts_none(kit, real_manager): + with real_manager() as mgr: + source = kit.published(mgr, SOURCE, PROMPT) + target = kit.admitted(mgr, TARGET, OTHER_PROMPT) + before = set(threading.enumerate()) + + def in_turn(call, name): + """``call`` on its own thread, joined; like the executor's, it used CUDA before.""" + + def run(): + torch.cuda.synchronize() + return call() + + return kit.on_thread(run, name) + + built = in_turn(lambda: attach(mgr, fetch_tokens=END, max_fetches=2), "builder") + lender = built["value"] + + def executor_loop(): + read = lender.lend_read(source, 0, END) + view = ready_view(read, mgr) + # A backend reads the arrays on its own thread until the release. + seen = kit.on_thread(lambda: [bytes(n) for n in view.runs[0].names], "backend") + write = lender.lend_write(target, 0, END) + write.mark_arrived(write.poll().row_masks(True)) + read.release() + write.release() + return seen["value"] + + looped = in_turn(executor_loop, "executor-loop") + assert "error" not in looped and len(looped["value"]) == BLOCKS + assert set(threading.enumerate()) <= before, "the lender started a thread" + assert "error" not in in_turn(mgr.shutdown, "shutdown") + with pytest.raises(ValueError): + lender.readiness(source) + + +# -- the manager's page-index buffer under a cache that outlives the manager -------------------- + +FREED = "touched the page-index buffer freed with its manager" + + +def check_a_staging_lease_outliving_its_manager(kit, attach): + """A read lease, which holds its request's cache, outlives its manager and lender, collected + without a shutdown or a free; a new tensor reclaims the freed buffer's address with a canary. + Dropping the lease closes the cache, which must write nothing there.""" + torch.cuda.init() + gc.collect() + torch.cuda.empty_cache() + mgr = kit.make_manager() + source = kit.published(mgr, SOURCE, PROMPT) + kv = kit.kv(mgr, source) + blocks = int(kv.num_blocks) + lender = attach(mgr, fetch_tokens=END) + lease = lender.lend_read(source, 0, END) + ready_view(lease, mgr) + row = kit.index_row(mgr, source) + assert row.values[:blocks] == list(kv.get_base_page_indices(0)[:blocks]) + watched = [weakref.ref(o) for o in (mgr.host_kv_cache_block_offsets, mgr, lender)] + mgr._stream.synchronize() + del mgr, lender + gc.collect() + assert watched[1]() is None and watched[2]() is None, "the manager or lender was not collected" + canary = None if watched[0]() is not None else kit.reclaim(row) + if watched[0]() is None: + assert canary is not None, "inconclusive: no allocation reused the freed buffer's address" + read = list(kv.get_base_page_indices(0)[:blocks]) + del kv + lease.release() + del lease # it held the cache's last reference + gc.collect() + written = kit.canary_written(canary, row) + assert not (canary is not None and read == [kit.CANARY] * blocks) and not written, ( + f"the leased cache {FREED}: it read {read} and its close wrote {written} (cell, value)" + ) + + +def test_a_lease_outliving_its_manager_touches_no_freed_index_buffer(kit): + check_a_staging_lease_outliving_its_manager(kit, attach) + + +def test_the_check_catches_a_lender_not_keeping_the_index_buffer(kit): + liar = attach_breaking({"_keep_index_buffer": lambda self, manager: None}) + with pytest.raises(CAUGHT, match=FREED): + check_a_staging_lease_outliving_its_manager(kit, liar) + + +def check_the_index_buffer_is_kept_until_the_shutdown(kit, real_manager, attach): + with real_manager() as mgr: + before = len(kit.retained()) + buffer = mgr.host_kv_cache_block_offsets + attach(mgr, fetch_tokens=END) + assert any(o is buffer for o in kit.retained()), "not kept from the attach" + mgr.shutdown() + assert not any(o is buffer for o in kit.retained()), "kept past the shutdown" + assert len(kit.retained()) == before + + +def test_the_index_buffer_is_kept_from_the_attach_until_the_shutdown(kit, real_manager): + check_the_index_buffer_is_kept_until_the_shutdown(kit, real_manager, attach) + + +def test_the_check_catches_a_lender_keeping_the_index_buffer_past_the_shutdown(kit, real_manager): + liar = attach_breaking({"_let_go_index_buffer": lambda self: None}) + with pytest.raises(CAUGHT, match="kept past the shutdown"): + check_the_index_buffer_is_kept_until_the_shutdown(kit, real_manager, liar) + + +# -- DeepSeek-V4 ------------------------------------------------------------------------------ + + +@skip_pre_blackwell +def test_a_deepseek_v4_prompt_is_fetched_through_staging_byte_for_byte(kit, deepseek_v4_manager): + tpb = 128 + end = 3 * tpb + prompt = list(range(3000, 3000 + end + 1)) + with deepseek_v4_manager() as mgr_a, deepseek_v4_manager() as mgr_b: + source = kit.published(mgr_a, SOURCE, prompt) + lender_a = attach(mgr_a, fetch_tokens=end) + publish = lender_a.lend_read(source, 0, end) + publish_view = ready_view(publish, mgr_a) + assert len(publish_view.runs) == kit.num_layer_groups(mgr_a) + assert check_staged(kit, mgr_a, lender_a, source, publish_view) == publish_view.num_rows + target = kit.admitted(mgr_b, TARGET, prompt) + lender_b = attach(mgr_b, fetch_tokens=end) + assert [p.name for p in lender_a.parts] == [p.name for p in lender_b.parts] + lease = lender_b.lend_write(target, 0, end) + view = lease.poll() + assert by_group(view) == by_group(publish_view) + kit.fill_sentinel(mgr_b, target) + lease.mark_arrived(kit.relay(lender_a, publish_view, lender_b, view)) + mgr_b._stream.synchronize() + assert lender_b.readiness(target) == (end, end) + for run in view.runs: + for ordinal in run.ordinals.tolist(): + lg = run.layer_group + got = kit.digest(kit.page(mgr_b, target, lg, ordinal)) + assert got == kit.digest(kit.page(mgr_a, source, lg, ordinal)) + lease.release() + publish.release() diff --git a/tests/unittest/_torch/executor/test_disagg_receive_ordering.py b/tests/unittest/_torch/executor/test_disagg_receive_ordering.py index d55a6c3fc232..e5b6ed12b523 100644 --- a/tests/unittest/_torch/executor/test_disagg_receive_ordering.py +++ b/tests/unittest/_torch/executor/test_disagg_receive_ordering.py @@ -50,6 +50,7 @@ def _coordinator(manager: Mock, receive: Mock) -> DisaggTransferCoordinator: def _manager() -> Mock: """Run the real admission and resource hooks with mocked cache allocation.""" manager = Mock(spec=KVCacheManagerV2) + manager._sharing = None # no lender: free_resources closes and removes the cache itself manager._stream = Mock() manager._disagg_receive_ready = {} manager.is_draft = False From 8562acc19139f40b6cb889d9bcf26f52e9333195 Mon Sep 17 00:00:00 2001 From: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> Date: Fri, 2 Oct 2026 16:32:42 -0700 Subject: [PATCH 2/2] [None][feat] RI-01 pin shared KV contracts and staging compatibility Signed-off-by: Chien-Chun Hung <2679986+chienchunhung@users.noreply.github.com> --- .../developer-guide/shared-kv-runtime.md | 64 +++ docs/source/index.rst | 1 + .../_torch/disaggregation/base/shared.py | 355 +++++++++++++++++ .../_torch/disaggregation/resource/shared.py | 210 ++++++++++ .../test_shared_backend_contract.py | 374 ++++++++++++++++++ .../test_shared_resource_contract.py | 334 ++++++++++++++++ 6 files changed, 1338 insertions(+) create mode 100644 docs/source/developer-guide/shared-kv-runtime.md create mode 100644 tensorrt_llm/_torch/disaggregation/base/shared.py create mode 100644 tensorrt_llm/_torch/disaggregation/resource/shared.py create mode 100644 tests/unittest/disaggregated/test_shared_backend_contract.py create mode 100644 tests/unittest/disaggregated/test_shared_resource_contract.py diff --git a/docs/source/developer-guide/shared-kv-runtime.md b/docs/source/developer-guide/shared-kv-runtime.md new file mode 100644 index 000000000000..2e659a393604 --- /dev/null +++ b/docs/source/developer-guide/shared-kv-runtime.md @@ -0,0 +1,64 @@ + + +# Shared KV runtime contracts + +The shared backend contract lives in +`tensorrt_llm._torch.disaggregation.base.shared`. Its 13 symbols describe opaque +content units, logical outcomes, backend access completion, routes and memory +registration. The existing paired transfer contract remains in `base.backend`; +paired-path convergence is a separate runtime milestone. Import from the intended +module explicitly because several type names overlap. + +## Compatibility boundary + +`resource.shared.SharedRuntimeProfile` records explicit runtime assembly facts. +Its contract revision is pinned to `147ed68276e7fb89d5de60002d4e793a78707a8c`; +its backend revision identifies the actual package or build under test and is a +separate value. Validate the profile before taking a staging hold, registering +memory or exposing an extent. + +The first supported shape is native NIXL, `KVCacheManagerV2`, BF16 MHA written by +TRTLLM in HND layout, TP=DP=PP=CP=1, manager-owned host staging and committed whole +blocks. Remapping, offload, compression, recurrent state, extra buffer roles and +retry remain outside that shape. An unknown extra feature is rejected as well. +These checks describe compatibility requirements; they neither enable runtime +integration nor establish deployment qualification. Exact adapter qualification +belongs to LC-MC-21 and runtime qualification to RI-09. + +The assembly caller must obtain the cache dtype, writer/layout and feature facts +from the actual engine configuration. A lender layout digest alone does not prove +compatibility: dtype and computational meaning belong to the reuse scope, and +the writer must really use the declared byte layout. No new capabilities method +is added to the shared backend API. + +## Content and physical identity + +Use the public `pyexecutor.kv_cache.sharing` lender types. `build_extent` copies +`GroupRun.names[i].tobytes()` unchanged into `Unit.name`; it never reconstructs a +content key. A unit's `local_group` is the lender's layer group and its `local` +coordinate is `(address - part.address) // part.slot_bytes`. Logical block +ordinals are not local storage slots. + +`STAGING_EXTENT_NAMESPACE` is a stable versioned adapter convention shared across +requests. Do not substitute request IDs or allocation identities. A `Part.name` +identifies a layout-compatible region, so matching names across workers do not +identify the same physical memory. Adapter, lease, allocation and registration +instances retain their distinct local ownership identities. + +## Delivery and readiness + +`Delivered.served` contains complete units only; an empty set is a miss. +`served_masks` rejects unknown names and maps the delivered subset into the +lender's row masks. It does not compute engine readiness or fill holes between +delivered units. On the manager's owner thread, call `Lease.mark_arrived` once, +then use `StagingLender.readiness` after the manager's local copy processing. +Across ranks, use the minimum `usable_until` and maximum `restart_floor`. + +Logical completion and backend access completion are independent. A failed or +cancelled outcome, request exit, and a timeout cannot authorize `Lease.release`. +Keep leases and registration roots until physical access and local copies have +ended; only then deregister memory and release the `PartsHold` on the manager's +owner thread. `release` asserts that access ended; it does not stop a transfer. diff --git a/docs/source/index.rst b/docs/source/index.rst index 4936eb2b70e1..9c0161c0ecd1 100644 --- a/docs/source/index.rst +++ b/docs/source/index.rst @@ -106,6 +106,7 @@ Welcome to TensorRT LLM's Documentation! developer-guide/dev-containers.md developer-guide/api-change.md developer-guide/kv-transfer.md + developer-guide/shared-kv-runtime.md developer-guide/kv-cache-cold-page-codec.md developer-guide/kv-cache-compression-development.md developer-guide/sparse-attention-development-guide.md diff --git a/tensorrt_llm/_torch/disaggregation/base/shared.py b/tensorrt_llm/_torch/disaggregation/base/shared.py new file mode 100644 index 000000000000..6bf34f3ce038 --- /dev/null +++ b/tensorrt_llm/_torch/disaggregation/base/shared.py @@ -0,0 +1,355 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Content-addressed cache backend contract. + +Shared-cache consumers import this module explicitly. The paired transceiver +uses ``base.backend`` and its package re-exports; its request-scoped extents and +outcomes are distinct from the process-lifetime backend contract here. + +Names encode content, reuse scope, layout, layer identity, and relevant sharding. +Backends treat them as opaque bytes; local coordinates resolve memory separately. +Logical outcomes never authorize memory reuse: only ``quiesce`` returning +``True`` proves that the listed deliveries have ended all caller-memory access. +These protocols describe obligations, not a mechanism that enforces them. +""" + +from __future__ import annotations + +from collections.abc import Iterable, Mapping, Sequence +from dataclasses import dataclass +from typing import Protocol, runtime_checkable + +__all__ = [ + "Unit", + "CacheExtent", + "Delivered", + "Failed", + "Cancelled", + "Outcome", + "Attempt", + "SubmissionRejected", + "Route", + "Registration", + "Fetches", + "Publishes", + "RegistersPools", +] + + +@dataclass(frozen=True, kw_only=True) +class Unit: + """One independently named, indivisibly delivered piece of cache content. + + Args: + name: Opaque identity, unique within the enclosing extent's name. + local_group: Nonnegative process-local layer-group ordinal. It must not + be sent to a peer or compared with a peer's group ordinal. + local: Nonnegative region identifier within that local group. Together + the two coordinates resolve one or more address-and-length spans. + """ + + name: bytes + local_group: int + local: int + + def __post_init__(self) -> None: + """Validate local coordinates. + + Raises: + ValueError: A coordinate is negative. + """ + if self.local_group < 0 or self.local < 0: + raise ValueError(f"negative local coordinate ({self.local_group}, {self.local})") + + +@dataclass(frozen=True, kw_only=True) +class CacheExtent: + """Immutable description of one fetch destination or publish source. + + Name derivation must include every quantity whose cross-side mismatch could + produce incorrect bytes. Layout, content, reuse scope, layer identity, token + coverage, and relevant sharding therefore belong in names, not extra fields. + This interface cannot validate that derivation. + + Args: + name: Opaque content name scoping all unit names in this delivery. + units: Units with distinct names. An empty extent is valid. Construction + snapshots the supplied sequence into a tuple. + is_last: Explicit sequence-end marker. Only sequence-aware backends may + inspect it; a name-addressed store must not use it. + """ + + name: bytes + units: tuple[Unit, ...] + is_last: bool + + def __post_init__(self) -> None: + """Snapshot the unit sequence and validate names within this extent. + + Raises: + ValueError: Two units have the same name. + """ + object.__setattr__(self, "units", tuple(self.units)) + names = [unit.name for unit in self.units] + if len(set(names)) != len(names): + raise ValueError("two units in one extent share a name") + + +@dataclass(frozen=True) +class Delivered: + """A logical answer listing only fully served units from this delivery. + + Empty ``served`` is a fetch miss or an unaccepted publication, not an error. + Unserved destinations must remain untouched. The caller correlates names to + retained coverage records and performs readiness checks before consumption; + this outcome establishes neither engine readiness nor physical quiescence. + + Args: + served: Subset of the requested unit names whose delivery completed. + """ + + served: frozenset[bytes] + + +@dataclass(frozen=True) +class Failed: + """Logical failure; destination contents are undefined until overwritten. + + A miss must not be reported as failure, nor a transport, capacity, or + registration failure as a miss. Failure may precede physical quiescence. + + Args: + reason: Human-readable failure description. + """ + + reason: str + + +@dataclass(frozen=True) +class Cancelled: + """Logical cancellation without delivery, initiated outside this interface. + + There is no cancellation operation on this interface. Submitted work may + continue accessing memory after cancellation is reported. + + Args: + by_peer: Whether the peer initiated cancellation. Peer cancellation is + a transfer error; local cancellation is ordinary termination. + """ + + by_peer: bool + + +Outcome = Delivered | Failed | Cancelled + + +@runtime_checkable +class Attempt(Protocol): + """One delivery with a stable logical outcome, independent of memory access.""" + + def poll(self) -> Outcome | None: + """Inspect the current logical result without blocking. + + Returns: + None while pending, otherwise an outcome that never changes on + subsequent polls. No outcome proves physical quiescence. + """ + ... + + +class SubmissionRejected(Exception): + """Submission rejected before work, peer notification, or memory access escaped. + + The caller has nothing to quiesce. After any such effect escapes, submission + must return an Attempt and report failure through its logical outcome. + """ + + +class Route(Protocol): + """Opaque per-request source plan created by a backend; it retains no content.""" + + def close(self) -> None: + """End route preparation after the caller stops submitting along it. + + Existing deliveries continue and still require quiescence. Successful + closure is idempotent. A raised exception leaves the route open and + retryable; it must not be recorded as successfully closed. + """ + ... + + +class Registration(Protocol): + """Opaque handle for one non-overlapping registered memory span.""" + + def close(self) -> None: + """Revoke registration after all deliveries using it pass quiescence. + + This operation neither waits for access to end nor proves memory safety. + Success is idempotent by handle: closing an old handle again must not + revoke a later registration of the same address. If closure raises, the + span remains registered and the same handle can be retried. + """ + ... + + +@runtime_checkable +class Fetches(Protocol): + """Process-lifetime fetching backend with all five methods required. + + Calls may overlap across requests, directions, and waits on the same + Attempt. Runtime protocol checks establish member presence, not signature, + concurrency, or semantic conformance. Observable per-unit hit counters are + required operationally but their form is outside this interface. + """ + + def fetch(self, extent: CacheExtent, *, route: Route | None = None) -> Attempt: + """Submit a fetch immediately without waiting for completion. + + Args: + extent: Content and local destination coordinates. + route: Unmodified open route produced by this backend, or None. + + Returns: + Handle for every submission whose effects escaped, including failed + submissions. For backends implementing RegistersPools, unregistered + destinations must yield Failed. + + Raises: + SubmissionRejected: Nothing escaped. This includes an unsupported, + foreign, or closed route rejected before any submission effect. + """ + ... + + def quiesce(self, attempts: Iterable[Attempt]) -> bool: + """Wait for a memory-access determination for only the listed deliveries. + + Args: + attempts: Submitted deliveries, possibly passed to this wait before. + + Returns: + True only when all listed deliveries will never again access caller + memory. False means this cannot be established, without promising + that waiting longer will help. Retire their spans without freeing, + reconstructing, or reusing them until a later call returns True. + Neither answer implies that a logical outcome exists. + """ + ... + + def settle(self, attempts: Iterable[Attempt]) -> None: + """Wait until every listed delivery has a stable logical outcome. + + This wait is unbounded and must not wait for unlisted deliveries. It + does not establish physical quiescence; use timed polling for a bounded + wait and quiesce for the separate memory-access question. + + Args: + attempts: Submitted deliveries, possibly passed to this wait before. + """ + ... + + def probe(self, name: bytes, units: Sequence[bytes]) -> frozenset[bytes] | None: + """Report advisory hits without moving payload or reserving content. + + Args: + name: Opaque content name scoping the requested units. + units: Requested unit names. + + Returns: + A subset of requested names, empty for no hits, or None if probing + cannot answer more cheaply than fetching. A later fetch may miss + previously reported hits. Probe failures must raise, never become + an empty set or None. + """ + ... + + def open_route(self, hint: Mapping[str, object]) -> Route: + """Prepare a per-request route without waiting or touching cache memory. + + Preparation may initiate control-plane work but must not move payload + or retain source content. A valid hint's transport error propagates + unchanged; this operation must never raise SubmissionRejected. + + Args: + hint: Deployment-defined routing information, opaque to this API. + + Returns: + An opaque route belonging to this backend and request. + + Raises: + NotImplementedError: This is a single-source backend. + ValueError: The hint is unrecognized, incomplete, or names an + unknown source. + """ + ... + + +@runtime_checkable +class Publishes(Protocol): + """Process-lifetime publication, independent of the Fetches capability. + + Concurrent calls and repeated waits must be supported. Publication is + atomically visible per unit, not per extent. Repeated publication under a + content name merges units without removing previously published units. + """ + + def publish(self, extent: CacheExtent) -> Attempt: + """Submit publication immediately using already-readable source memory. + + Args: + extent: Content and local source coordinates. No route is supplied: + the backend is the store or answers the requesting peer. + + Returns: + Handle for every submission whose effects escaped, including failed + submissions. For backends implementing RegistersPools, unregistered + sources must yield Failed. + + Raises: + SubmissionRejected: No work, notification, or memory access escaped. + """ + ... + + def quiesce(self, attempts: Iterable[Attempt]) -> bool: + """Establish physical quiescence with the same semantics as Fetches. + + Args: + attempts: Deliveries whose memory access must end; unlisted + deliveries must not determine when this wait returns. + + Returns: + True only if caller-memory access has ended permanently. False + requires retaining all affected memory and promises no eventual + True answer. Logical outcomes remain independent. + """ + ... + + def settle(self, attempts: Iterable[Attempt]) -> None: + """Wait for logical outcomes with the same semantics as Fetches. + + Args: + attempts: Deliveries that must have stable outcomes before return. + Unlisted deliveries must not determine when this wait returns. + """ + ... + + +@runtime_checkable +class RegistersPools(Protocol): + """Optional pool registration, independent of local-coordinate resolution.""" + + def register_pool(self, address: int, size: int) -> Registration: + """Register one memory span before any delivery refers to it. + + Register each span once; overlapping live registrations must be + rejected. Failure raises and must register nothing. Registration does + not establish how local coordinates resolve to this address interval. + + Args: + address: Starting memory address accessible to the backend. + size: Span length in bytes. + + Returns: + A handle that revokes exactly this registration after all referring + deliveries are quiescent and new submissions have stopped. + """ + ... diff --git a/tensorrt_llm/_torch/disaggregation/resource/shared.py b/tensorrt_llm/_torch/disaggregation/resource/shared.py new file mode 100644 index 000000000000..1d1810c508b6 --- /dev/null +++ b/tensorrt_llm/_torch/disaggregation/resource/shared.py @@ -0,0 +1,210 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Map public manager staging views to the shared backend contract.""" + +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass +from numbers import Integral + +import numpy as np + +from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import Part, RegionView + +from ..base.shared import CacheExtent, Unit + +SHARED_CONTRACT_REVISION = "147ed68276e7fb89d5de60002d4e793a78707a8c" +STAGING_EXTENT_NAMESPACE = b"trtllm:shared-kv:staging:1" + + +@dataclass(frozen=True) +class SharedRuntimeProfile: + """Explicit assembly facts for the first shared runtime integration profile. + + Args: + contract_revision: Exact shared contract revision used by the adapter. + backend_revision: Concrete backend package/build revision for diagnostics. + backend_kind: Backend implementation family. + manager_kind: Manager implementation family. + cache_dtype: Element type of the KV cache. + attention_backend: Backend that actually writes the cache. + attention_kind: Attention cache kind. + layout: Physical layout actually written by the attention backend. + parallelism: Tensor, data, pipeline and context parallel sizes, in order. + staging: Owner and location of transfer staging memory. + committed_whole_blocks: Whether content consists of committed whole blocks. + extra_features: Features beyond the supported profile, including remapping, + offload, compression, recurrent state, extra buffer roles or retries. + + These facts are supplied by runtime assembly, not inferred from a layout + digest. Validation restricts compatibility; it does not certify a backend or + enable a scheduler path. The backend revision is independent of the contract + revision and must identify the actual implementation under qualification. + """ + + contract_revision: str + backend_revision: str + backend_kind: str + manager_kind: str + cache_dtype: str + attention_backend: str + attention_kind: str + layout: str + parallelism: tuple[int, int, int, int] + staging: str + committed_whole_blocks: bool + extra_features: frozenset[str] + + def validate(self) -> None: + """Reject unsupported assembly facts before acquiring or exposing resources. + + Raises: + ValueError: The revision is missing or the profile is unsupported. + """ + expected = { + "contract_revision": (self.contract_revision, SHARED_CONTRACT_REVISION), + "backend_kind": (self.backend_kind, "native_nixl"), + "manager_kind": (self.manager_kind, "KVCacheManagerV2"), + "cache_dtype": (self.cache_dtype, "bfloat16"), + "attention_backend": (self.attention_backend, "TRTLLM"), + "attention_kind": (self.attention_kind, "mha"), + "layout": (self.layout, "HND"), + "staging": (self.staging, "manager_host"), + } + for field, (actual, supported) in expected.items(): + if actual != supported: + raise ValueError(f"unsupported shared runtime {field}: {actual!r}") + if not isinstance(self.backend_revision, str) or not self.backend_revision.strip(): + raise ValueError("backend_revision must identify a concrete backend build") + if ( + not isinstance(self.parallelism, tuple) + or len(self.parallelism) != 4 + or any(type(size) is not int or size != 1 for size in self.parallelism) + ): + raise ValueError("shared runtime requires TP=DP=PP=CP=1") + if self.committed_whole_blocks is not True: + raise ValueError("shared runtime requires committed whole blocks") + if not isinstance(self.extra_features, frozenset) or self.extra_features: + raise ValueError(f"unsupported shared runtime extra_features: {self.extra_features!r}") + + +def _check_parts(parts: Sequence[Part]) -> None: + """Validate host-region geometry before deriving slot coordinates. + + Args: + parts: Public staging allocation descriptions. + + Raises: + ValueError: A region is malformed or overlaps another region. + """ + spans = [] + for part in parts: + for value in (part.address, part.nbytes, part.slot_bytes, part.slots): + if isinstance(value, bool) or not isinstance(value, Integral) or value <= 0: + raise ValueError("staging part geometry must contain positive integers") + if part.slots * part.slot_bytes > part.nbytes: + raise ValueError("staging part does not cover its slots") + spans.append((int(part.address), int(part.address) + int(part.nbytes))) + spans.sort() + if any(left[1] > right[0] for left, right in zip(spans, spans[1:])): + raise ValueError("staging parts overlap") + + +def build_extent( + view: RegionView, parts: Sequence[Part], *, name: bytes, is_last: bool +) -> CacheExtent: + """Translate a ready staging lease to whole opaque content units. + + Args: + view: Ready view returned by a public staging lease. + parts: The same lender's fixed staging regions. + name: Stable versioned namespace, normally ``STAGING_EXTENT_NAMESPACE``; + never a request ID or physical allocation identity. + is_last: Whether this extent ends the logical operation. + + Returns: + An immutable extent preserving every lender row name byte for byte. + Unit coordinates identify local staging slots, not token block ordinals. + + Raises: + ValueError: Metadata is absent, out of bounds, misaligned or duplicated. + TypeError: The extent namespace or final-extent flag is invalid. + """ + if not isinstance(name, bytes): + raise TypeError("extent namespace must be bytes") + if not isinstance(is_last, bool): + raise TypeError("is_last must be bool") + _check_parts(parts) + units = [] + occupied = set() + for run in view.runs: + if run.names is None or run.addresses is None or run.part is None: + raise ValueError("shared extents require staging names, addresses and part indices") + if run.part >= len(parts): + raise ValueError("staging run refers to an unknown part") + if np.any(run.ordinals < 0): + raise ValueError("staging block ordinals must be nonnegative") + part = parts[run.part] + for index, address in enumerate(run.addresses): + offset = int(address) - int(part.address) + if offset < 0 or offset % part.slot_bytes: + raise ValueError("staging row is outside or misaligned with its part") + slot = offset // int(part.slot_bytes) + if slot >= part.slots or offset + part.slot_bytes > part.nbytes: + raise ValueError("staging row exceeds its part capacity") + coordinate = (run.part, slot) + if coordinate in occupied: + raise ValueError("staging rows contain a duplicate physical coordinate") + occupied.add(coordinate) + units.append( + Unit(name=run.names[index].tobytes(), local_group=run.layer_group, local=slot) + ) + return CacheExtent(name=name, units=tuple(units), is_last=is_last) + + +def served_masks(view: RegionView, served: frozenset[bytes]) -> tuple[np.ndarray, ...]: + """Convert delivered whole-unit names to the lender's arrival masks. + + Args: + view: The ready staging view used to construct the submitted extent. + served: Names reported by ``Delivered.served``; an empty set is a miss. + + Returns: + One boolean row mask per run, suitable for ``Lease.mark_arrived``. + These masks describe delivery only; readiness still comes from the lender + after local copies and contiguous-prefix checks. + + Raises: + TypeError: ``served`` is not a frozen set of byte names. + ValueError: A view lacks names, duplicates names, or a served name was + absent from the submitted view. + """ + if not isinstance(served, frozenset) or any(not isinstance(name, bytes) for name in served): + raise TypeError("served must be a frozenset of opaque byte names") + masks = view.row_masks() + known = set() + for run, mask in zip(view.runs, masks): + if run.names is None: + raise ValueError("arrival masks require staging row names") + for index, row in enumerate(run.names): + name = row.tobytes() + if name in known: + raise ValueError("staging view contains duplicate names") + known.add(name) + mask[index] = name in served + if not served <= known: + raise ValueError("served names must be a subset of the submitted extent") + return masks diff --git a/tests/unittest/disaggregated/test_shared_backend_contract.py b/tests/unittest/disaggregated/test_shared_backend_contract.py new file mode 100644 index 000000000000..f5d9f9b28081 --- /dev/null +++ b/tests/unittest/disaggregated/test_shared_backend_contract.py @@ -0,0 +1,374 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""CPU tests for the shared backend data and protocol surface. + +The controllable backend supplies examples of independent logical and physical +completion. Those examples validate the test double, not provider conformance; +each real provider must establish these guarantees with its own transfer tests. +""" + +from collections.abc import Iterable, Mapping, Sequence +from dataclasses import FrozenInstanceError, fields +from inspect import Parameter, signature +from threading import Event +from types import SimpleNamespace +from typing import get_type_hints + +import pytest + +from tensorrt_llm._torch.disaggregation.base import shared +from tensorrt_llm._torch.disaggregation.base.shared import ( + Attempt, + CacheExtent, + Cancelled, + Delivered, + Failed, + Fetches, + Outcome, + Publishes, + RegistersPools, + Registration, + Route, + SubmissionRejected, + Unit, +) + +pytestmark = pytest.mark.cpu_only + + +class _ControlledAttempt: + """An outcome and physical-access evidence advanced separately by a test.""" + + def __init__(self) -> None: + """Initialize a pending delivery with no evidence of ended access.""" + self._outcome: Outcome | None = None + self.logical_ready = Event() + self.access_ended = Event() + + def poll(self) -> Outcome | None: + """Return the immutable logical outcome, if supplied by the test.""" + return self._outcome + + def answer(self, outcome: Outcome) -> None: + """Commit one logical answer without changing access evidence. + + Args: + outcome: Answer to expose on every subsequent poll. + + Raises: + ValueError: An answer has already been committed. + """ + if self._outcome is not None: + raise ValueError("attempt already has an outcome") + self._outcome = outcome + self.logical_ready.set() + + +class _ControlledStore: + """Single-source fetch double backed by local bytearrays and explicit events.""" + + def __init__(self, memory: dict[tuple[int, int], bytearray]) -> None: + """Initialize a double without populating remote content. + + Args: + memory: Destination spans keyed by local group and region. + """ + self.memory = memory + self.content: dict[tuple[bytes, bytes], bytes] = {} + self.attempts: list[_ControlledAttempt] = [] + self.extents: list[CacheExtent] = [] + self.hits = 0 + self.reject = False + self.fail_after_write = False + + def fetch(self, extent: CacheExtent, *, route: Route | None = None) -> _ControlledAttempt: + """Record submission, with controlled rejection or escaped failure. + + Args: + extent: Names and destination coordinates. + route: Unsupported by this single-source double. + + Returns: + An attempt for every accepted submission. + + Raises: + SubmissionRejected: Rejection was requested or a route was supplied; + no destination, counter, or submitted-attempt state changes. + """ + if self.reject or route is not None: + raise SubmissionRejected("submission rejected before effects") + attempt = _ControlledAttempt() + self.attempts.append(attempt) + self.extents.append(extent) + if self.fail_after_write: + unit = extent.units[0] + self.memory[unit.local_group, unit.local][0] = 0 + attempt.answer(Failed("transport failed after touching destination")) + return attempt + + def deliver(self, attempt: _ControlledAttempt) -> None: + """Write complete hits and provide an answer without ending access. + + Args: + attempt: Delivery previously submitted to this double. + """ + extent = self.extents[self.attempts.index(attempt)] + served: set[bytes] = set() + for unit in extent.units: + payload = self.content.get((extent.name, unit.name)) + if payload is not None: + destination = self.memory[unit.local_group, unit.local] + if len(destination) != len(payload): + attempt.answer(Failed("matching name has incompatible destination")) + return + destination[:] = payload + served.add(unit.name) + self.hits += 1 + attempt.answer(Delivered(frozenset(served))) + + def quiesce(self, attempts: Iterable[Attempt]) -> bool: + """Report only explicit no-future-access evidence supplied by the test. + + Args: + attempts: Deliveries to inspect independently of logical outcomes. + + Returns: + True if every delivery has ended access; False otherwise, with no + guarantee that evidence will ever be supplied. + """ + return all( + self.attempts[self.attempts.index(attempt)].access_ended.is_set() + for attempt in attempts + ) + + def settle(self, attempts: Iterable[Attempt]) -> None: + """Wait only for the supplied deliveries' logical outcomes. + + Args: + attempts: Deliveries whose outcomes must exist on return. + """ + for attempt in attempts: + self.attempts[self.attempts.index(attempt)].logical_ready.wait() + + def probe(self, name: bytes, units: Sequence[bytes]) -> frozenset[bytes]: + """Return held unit names without reserving them or touching destinations. + + Args: + name: Extent name scoping the requested unit names. + units: Candidate unit names. + + Returns: + The requested unit names currently held by this double. + """ + return frozenset(unit for unit in units if (name, unit) in self.content) + + def open_route(self, hint: Mapping[str, object]) -> Route: + """Reject routing for this single-source backend. + + Args: + hint: Unused deployment hint. + + Raises: + NotImplementedError: Single-source backends have no route to open. + """ + raise NotImplementedError("single-source backend") + + +def _extent() -> CacheExtent: + """Return two units with equal region IDs in different local groups.""" + return CacheExtent( + name=b"content", + units=( + Unit(name=b"first", local_group=0, local=0), + Unit(name=b"second", local_group=1, local=0), + ), + is_last=True, + ) + + +def test_export_surface_and_fields_match_canonical_contract() -> None: + """Keep incompatible content metadata out of the backend's field surface.""" + assert set(shared.__all__) == { + "Unit", + "CacheExtent", + "Delivered", + "Failed", + "Cancelled", + "Outcome", + "Attempt", + "SubmissionRejected", + "Route", + "Registration", + "Fetches", + "Publishes", + "RegistersPools", + } + assert len(shared.__all__) == 13 + for cls, names in ( + (Unit, ["name", "local_group", "local"]), + (CacheExtent, ["name", "units", "is_last"]), + (Delivered, ["served"]), + (Failed, ["reason"]), + (Cancelled, ["by_peer"]), + ): + assert [field.name for field in fields(cls)] == names + assert set(get_type_hints(Fetches.quiesce)) == {"attempts", "return"} + assert get_type_hints(Fetches.quiesce)["return"] is bool + assert get_type_hints(Publishes.quiesce)["return"] is bool + route = signature(Fetches.fetch).parameters["route"] + assert route.kind is Parameter.KEYWORD_ONLY and route.default is None + assert set(signature(Publishes.publish).parameters) == {"self", "extent"} + + +@pytest.mark.parametrize("local_group,local", [(-1, 0), (0, -1), (-1, -1)]) +def test_negative_coordinates_are_rejected(local_group: int, local: int) -> None: + """A unit must resolve through nonnegative local coordinates.""" + with pytest.raises(ValueError, match="negative local coordinate"): + Unit(name=b"opaque", local_group=local_group, local=local) + + +def test_extent_snapshots_sequence_and_rejects_duplicate_names() -> None: + """Caller list mutation cannot change a lent extent or introduce aliases.""" + units = list(_extent().units) + extent = CacheExtent(name=b"content", units=units, is_last=False) + units.clear() + assert extent.units == _extent().units + with pytest.raises(ValueError, match="share a name"): + CacheExtent( + name=b"content", + units=(extent.units[0], Unit(name=b"first", local_group=9, local=7)), + is_last=True, + ) + + +@pytest.mark.parametrize( + "value,attribute,new_value", + [ + (Unit(name=b"first", local_group=0, local=0), "local", 1), + (_extent(), "units", ()), + (Delivered(frozenset({b"first"})), "served", frozenset()), + (Failed("failure"), "reason", "replacement"), + (Cancelled(False), "by_peer", True), + ], +) +def test_descriptions_and_outcomes_are_frozen( + value: object, attribute: str, new_value: object +) -> None: + """Delivery descriptors and correctly typed outcomes reject reassignment.""" + with pytest.raises(FrozenInstanceError): + setattr(value, attribute, new_value) + + +def test_empty_extent_is_valid_and_markers_are_explicit() -> None: + """An empty delivery is distinct from an omitted sequence or cancel marker.""" + assert CacheExtent(name=b"", units=(), is_last=True).units == () + assert signature(CacheExtent).parameters["is_last"].default is Parameter.empty + assert signature(Cancelled).parameters["by_peer"].default is Parameter.empty + for parameter in signature(Unit).parameters.values(): + assert parameter.kind is Parameter.KEYWORD_ONLY + + +@pytest.mark.parametrize("missing", ["fetch", "quiesce", "settle", "probe", "open_route"]) +def test_fetch_requires_every_member(missing: str) -> None: + """Unavailable probing and routing are answers, not omitted capabilities.""" + store = _ControlledStore({}) + members = { + name: getattr(store, name) + for name in ("fetch", "quiesce", "settle", "probe", "open_route") + if name != missing + } + assert not isinstance(SimpleNamespace(**members), Fetches) + assert isinstance(store, Fetches) + assert not isinstance(store, Publishes) + assert not isinstance(store, RegistersPools) + assert isinstance(_ControlledAttempt(), Attempt) + + +def test_handles_do_not_advertise_runtime_protocol_checks() -> None: + """Opaque handles are returned unchanged; structural checks are not promised.""" + for protocol in (Route, Registration): + with pytest.raises(TypeError, match="runtime_checkable"): + isinstance(object(), protocol) + + +@pytest.mark.parametrize( + "hits", + [frozenset(), frozenset({b"first"}), frozenset({b"first", b"second"})], + ids=["miss", "partial", "complete"], +) +def test_controlled_store_leaves_unserved_destinations_untouched(hits: frozenset[bytes]) -> None: + """The fixture expresses whole-unit answers with distinct local group lookup.""" + extent = _extent() + store = _ControlledStore({(0, 0): bytearray(b"----"), (1, 0): bytearray(b"----")}) + store.content = {(extent.name, name): b"data" for name in hits} + attempt = store.fetch(extent) + assert attempt.poll() is None + store.deliver(attempt) + outcome = attempt.poll() + assert isinstance(outcome, Delivered) and outcome.served == hits + for unit in extent.units: + assert store.memory[unit.local_group, unit.local] == ( + b"data" if unit.name in hits else b"----" + ) + assert store.hits == len(hits) + assert not store.quiesce([attempt]) + + +def test_controlled_probe_is_advisory_and_single_source_routes_are_rejected() -> None: + """A hit may disappear after probing; rejected routes never submit a fetch.""" + store = _ControlledStore({(0, 0): bytearray(b"----"), (1, 0): bytearray(b"----")}) + store.content[b"content", b"first"] = b"data" + assert store.probe(b"content", [b"first", b"second"]) == frozenset({b"first"}) + store.content.clear() + attempt = store.fetch(_extent()) + store.deliver(attempt) + assert attempt.poll() == Delivered(frozenset()) + with pytest.raises(NotImplementedError): + store.open_route({"source": "worker"}) + assert len(store.attempts) == 1 + + +@pytest.mark.parametrize("escaped", [False, True], ids=["rejected", "escaped-failure"]) +def test_controlled_submission_retains_a_handle_after_effects_escape(escaped: bool) -> None: + """The fixture distinguishes no-effect rejection from a partially written failure.""" + store = _ControlledStore({(0, 0): bytearray(b"----"), (1, 0): bytearray(b"----")}) + store.reject, store.fail_after_write = not escaped, escaped + if not escaped: + with pytest.raises(SubmissionRejected): + store.fetch(_extent()) + assert store.attempts == [] and store.hits == 0 + assert store.memory == {(0, 0): bytearray(b"----"), (1, 0): bytearray(b"----")} + else: + attempt = store.fetch(_extent()) + assert isinstance(attempt.poll(), Failed) + assert store.memory[0, 0][0] == 0 + store.settle([attempt]) + assert not store.quiesce([attempt]) + attempt.access_ended.set() + assert store.quiesce([attempt]) + + +@pytest.mark.parametrize("physical_first", [False, True], ids=["outcome-first", "access-first"]) +def test_controlled_outcome_and_access_end_are_independent(physical_first: bool) -> None: + """Either event may precede the other, independently of unrelated attempts.""" + store = _ControlledStore({}) + extent = CacheExtent(name=b"empty", units=(), is_last=True) + attempt, unrelated = store.fetch(extent), store.fetch(extent) + outcome = Cancelled(by_peer=False) + if physical_first: + attempt.access_ended.set() + assert store.quiesce([attempt]) + assert attempt.poll() is None + attempt.answer(outcome) + else: + attempt.answer(outcome) + store.settle([attempt]) + assert not store.quiesce([attempt]) + attempt.access_ended.set() + for _ in range(2): + store.settle([attempt]) + assert store.quiesce([attempt]) + assert attempt.poll() is outcome + assert unrelated.poll() is None and not store.quiesce([unrelated]) + with pytest.raises(ValueError, match="already has an outcome"): + attempt.answer(Failed("late failure")) diff --git a/tests/unittest/disaggregated/test_shared_resource_contract.py b/tests/unittest/disaggregated/test_shared_resource_contract.py new file mode 100644 index 000000000000..7de31c9fd3e3 --- /dev/null +++ b/tests/unittest/disaggregated/test_shared_resource_contract.py @@ -0,0 +1,334 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""CPU contracts for actual lender-to-backend mapping and profile checks.""" + +from dataclasses import replace + +import numpy as np +import pytest + +from tensorrt_llm._torch.disaggregation.resource.shared import ( + SHARED_CONTRACT_REVISION, + STAGING_EXTENT_NAMESPACE, + SharedRuntimeProfile, + build_extent, + served_masks, +) +from tensorrt_llm._torch.pyexecutor.kv_cache.sharing import GroupRun, Part, RegionView + +pytestmark = pytest.mark.cpu_only + +NAME_A = bytes(range(54)) +NAME_B = b"\xff\x00" * 27 +NAME_C = b"\x80" * 54 +PART = Part(name="same-layout", address=4096, nbytes=1024, slot_bytes=128, slots=8) + + +def _view( + *, + names: tuple[bytes, ...] = (NAME_A, NAME_B), + addresses: tuple[int, ...] = (4480, 4736), + ordinals: tuple[int, ...] = (40, 41), + part: int = 0, + layer_group: int = 7, +) -> RegionView: + """Construct real public lender values with independently chosen coordinates. + + Args: + names: Opaque content names, including zero and non-ASCII bytes. + addresses: Host addresses of staging slots. + ordinals: Logical block positions, deliberately unlike slot positions. + part: Index of the containing registered region. + layer_group: Lender-local layer group index. + + Returns: + A public staging view for the mapping functions under test. + """ + return RegionView( + ( + GroupRun( + layer_group=layer_group, + ordinals=np.array(ordinals), + names=np.array([list(name) for name in names], dtype=np.uint8), + addresses=np.array(addresses), + part=part, + ), + ) + ) + + +def _profile() -> SharedRuntimeProfile: + """Return explicit first-profile assembly facts with a distinct backend revision. + + Returns: + A compatibility profile, which is not a deployment qualification record. + """ + return SharedRuntimeProfile( + contract_revision=SHARED_CONTRACT_REVISION, + backend_revision="native-nixl-test-build-42", + backend_kind="native_nixl", + manager_kind="KVCacheManagerV2", + cache_dtype="bfloat16", + attention_backend="TRTLLM", + attention_kind="mha", + layout="HND", + parallelism=(1, 1, 1, 1), + staging="manager_host", + committed_whole_blocks=True, + extra_features=frozenset(), + ) + + +def test_golden_names_and_slot_coordinates() -> None: + """Preserve opaque bytes and map addresses, not logical ordinals, into slots.""" + extent = build_extent(_view(), (PART,), name=STAGING_EXTENT_NAMESPACE, is_last=True) + + assert extent.name == b"trtllm:shared-kv:staging:1" + assert extent.is_last is True + assert [(unit.name, unit.local_group, unit.local) for unit in extent.units] == [ + (bytes(range(54)), 7, 3), + (b"\xff\x00" * 27, 7, 5), + ] + + +def test_relocation_changes_slots_without_renaming_content() -> None: + """Content identity is independent of allocation address and staging position.""" + before = build_extent(_view(), (PART,), name=STAGING_EXTENT_NAMESPACE, is_last=False) + relocated = replace(PART, address=8192) + after = build_extent( + _view(addresses=(8704, 8960)), + (relocated,), + name=STAGING_EXTENT_NAMESPACE, + is_last=False, + ) + + assert [unit.name for unit in before.units] == [unit.name for unit in after.units] + assert [unit.local for unit in after.units] == [4, 6] + assert before.name == after.name + assert PART.name == relocated.name + + +def test_multiple_groups_keep_local_group_identity() -> None: + """Local group numbers survive mapping even when groups share a part.""" + first = _view() + second = _view(names=(NAME_C,), addresses=(4992,), ordinals=(9,), layer_group=2) + extent = build_extent( + RegionView(first.runs + second.runs), (PART,), name=STAGING_EXTENT_NAMESPACE, is_last=True + ) + + assert [(unit.local_group, unit.local) for unit in extent.units] == [(7, 3), (7, 5), (2, 7)] + + +@pytest.mark.parametrize( + ("address", "error"), + [(3968, "outside"), (4481, "misaligned"), (5120, "capacity")], + ids=["before-part", "unaligned", "past-final-slot"], +) +def test_invalid_row_address_is_rejected(address: int, error: str) -> None: + """Do not expose a row whose complete bytes are outside a staging slot. + + Args: + address: Invalid host row address. + error: Expected diagnostic category. + """ + with pytest.raises(ValueError, match=error): + build_extent( + _view(addresses=(address, 4736)), + (PART,), + name=STAGING_EXTENT_NAMESPACE, + is_last=True, + ) + + +@pytest.mark.parametrize( + "part", + [ + replace(PART, address=0), + replace(PART, slot_bytes=0), + replace(PART, slots=True), + replace(PART, nbytes=1023), + ], + ids=["null-region", "zero-width", "boolean-capacity", "incomplete-last-slot"], +) +def test_malformed_part_is_rejected(part: Part) -> None: + """Reject malformed public allocation descriptions before using addresses. + + Args: + part: Invalid public region metadata. + """ + with pytest.raises(ValueError, match="staging part"): + build_extent(_view(), (part,), name=STAGING_EXTENT_NAMESPACE, is_last=True) + + +def test_overlapping_parts_are_rejected() -> None: + """Different region indices cannot disguise aliased physical storage.""" + with pytest.raises(ValueError, match="overlap"): + build_extent( + _view(), + (PART, replace(PART, address=4608)), + name=STAGING_EXTENT_NAMESPACE, + is_last=True, + ) + + +def test_unknown_part_is_rejected() -> None: + """A run must address a region belonging to the same lender.""" + with pytest.raises(ValueError, match="unknown part"): + build_extent(_view(part=1), (PART,), name=STAGING_EXTENT_NAMESPACE, is_last=True) + + +def test_in_place_view_is_rejected() -> None: + """Device-page views without staging metadata are outside this adapter.""" + view = RegionView((GroupRun(layer_group=0, ordinals=np.array([1])),)) + with pytest.raises(ValueError, match="staging names"): + build_extent(view, (PART,), name=STAGING_EXTENT_NAMESPACE, is_last=True) + + +def test_negative_ordinal_is_rejected() -> None: + """Only actual block rows can be mapped to whole-unit transfers.""" + with pytest.raises(ValueError, match="nonnegative"): + build_extent(_view(ordinals=(-1, 41)), (PART,), name=STAGING_EXTENT_NAMESPACE, is_last=True) + + +def test_duplicate_physical_coordinate_across_groups_is_rejected() -> None: + """Distinct content names must not alias one physical destination slot.""" + view = RegionView( + _view().runs + _view(names=(NAME_C,), addresses=(4480,), ordinals=(9,), layer_group=2).runs + ) + with pytest.raises(ValueError, match="duplicate physical coordinate"): + build_extent(view, (PART,), name=STAGING_EXTENT_NAMESPACE, is_last=True) + + +def test_duplicate_content_name_uses_canonical_extent_validation() -> None: + """One backend extent cannot ambiguously address the same content twice.""" + with pytest.raises(ValueError, match="share a name"): + build_extent( + _view(names=(NAME_A, NAME_A)), (PART,), name=STAGING_EXTENT_NAMESPACE, is_last=True + ) + + +@pytest.mark.parametrize( + ("served", "expected"), + [ + (frozenset(), [False, False]), + (frozenset({NAME_B}), [False, True]), + (frozenset({NAME_A, NAME_B}), [True, True]), + ], + ids=["miss", "partial-with-prefix-hole", "all-units"], +) +def test_served_masks_preserve_whole_rows(served: frozenset[bytes], expected: list[bool]) -> None: + """Map delivery exactly, without treating the highest delivered block as ready. + + Args: + served: Whole-unit backend result. + expected: Expected arrival mask in original lender row order. + """ + view = _view() + masks = served_masks(view, served) + + assert len(masks) == 1 + np.testing.assert_array_equal(masks[0], expected) + assert masks[0].dtype == np.bool_ + assert masks[0].flags.writeable + np.testing.assert_array_equal(view.runs[0].ordinals, [40, 41]) + + +def test_unknown_served_name_is_rejected() -> None: + """A backend cannot claim delivery of a unit absent from the submitted view.""" + with pytest.raises(ValueError, match="subset"): + served_masks(_view(), frozenset({NAME_C})) + + +def test_duplicate_view_names_cannot_mark_multiple_rows() -> None: + """Reject ambiguous arrival mapping even when called without extent construction.""" + with pytest.raises(ValueError, match="duplicate names"): + served_masks(_view(names=(NAME_A, NAME_A)), frozenset({NAME_A})) + + +def test_empty_view_and_miss_are_valid() -> None: + """Empty extents have no addresses to expose and no delivered rows.""" + view = RegionView(()) + assert build_extent(view, (), name=STAGING_EXTENT_NAMESPACE, is_last=True).units == () + assert served_masks(view, frozenset()) == () + + +def test_supported_profile_and_backend_revision_are_distinct() -> None: + """Accept explicit compatible facts without equating code and contract revisions.""" + profile = _profile() + profile.validate() + assert profile.backend_revision != profile.contract_revision + + +@pytest.mark.parametrize( + "changes", + [ + {"contract_revision": "unreviewed-contract"}, + {"backend_revision": " "}, + {"backend_kind": "mooncake"}, + {"manager_kind": "KVCacheManager"}, + {"cache_dtype": "float16"}, + {"attention_backend": "FLASHINFER"}, + {"attention_kind": "mla"}, + {"layout": "NHD"}, + {"parallelism": (2, 1, 1, 1)}, + {"parallelism": (1, 2, 1, 1)}, + {"parallelism": (1, 1, 2, 1)}, + {"parallelism": (1, 1, 1, 2)}, + {"staging": "python_bounce"}, + {"committed_whole_blocks": False}, + {"extra_features": frozenset({"compression"})}, + {"extra_features": frozenset({"retry"})}, + {"extra_features": frozenset({"unknown-feature"})}, + ], + ids=[ + "contract-revision", + "backend-revision", + "backend", + "manager", + "dtype", + "writer", + "cache-kind", + "layout", + "tp", + "dp", + "pp", + "cp", + "staging", + "partial-blocks", + "compression", + "retry", + "unknown-feature", + ], +) +def test_unsupported_profile_is_rejected(changes: dict[str, object]) -> None: + """Fail closed for profiles requiring separate compatibility and lifecycle work. + + Args: + changes: Unsupported facts replacing an otherwise supported profile. + """ + with pytest.raises(ValueError): + replace(_profile(), **changes).validate() + + +def test_request_id_is_not_an_extent_namespace() -> None: + """Reject a request-shaped integer where a shared namespace is required.""" + with pytest.raises(TypeError, match="namespace must be bytes"): + build_extent(_view(), (PART,), name=123, is_last=True) + + +def test_nonboolean_final_extent_flag_is_rejected() -> None: + """Do not silently reinterpret truthy objects as the final-extent flag.""" + with pytest.raises(TypeError, match="is_last must be bool"): + build_extent(_view(), (PART,), name=STAGING_EXTENT_NAMESPACE, is_last=1)