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/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 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)