From 68f0936314b38f06ca58e6a28036234538bc3c1e Mon Sep 17 00:00:00 2001 From: Yueh-Ting Chen Date: Thu, 24 Sep 2026 00:19:11 +0800 Subject: [PATCH] [TRTLLM-12891][feat] Budget connector prefixes during V2 scheduling Reserve stable source KV before admission and dispatch confirmed loads only after the final batch and destination allocations are ready. Release rejected promises and retain load ownership until every worker completes, including when the client cancels. Preserve the legacy connector query path. Cover admission credit, source protection, allocation rejection, replay, and cancellation draining with unit, real-manager, and integration tests. Signed-off-by: Yueh-Ting Chen --- docs/source/features/kv-cache-connector.md | 89 +++- examples/llm-api/llm_kv_cache_connector.py | 219 ++++++---- .../connectors/kv_cache_connector.py | 391 ++++++++++++++++- .../connectors/prefix_load_completion.py | 101 +++++ .../pyexecutor/executor_request_queue.py | 12 +- .../kv_cache/kv_cache_manager_v2.py | 145 +++++-- tensorrt_llm/_torch/pyexecutor/llm_request.py | 12 +- tensorrt_llm/_torch/pyexecutor/py_executor.py | 133 +++++- .../pyexecutor/scheduler/scheduler_v2.py | 28 +- .../defs/llmapi/test_llm_api_connector.py | 263 +++++++++++- .../integration/test_lists/test-db/l0_a10.yml | 5 + .../kv_cache/test_kv_cache_v2_scheduler.py | 108 +++++ .../test_kv_connector_executor_lifetime.py | 393 ++++++++++++++++++ .../test_kv_connector_reservations.py | 380 +++++++++++++++++ .../executor/test_kv_connector_v2_prefix.py | 265 +++++++++++- ...est_kv_connector_v2_prefix_real_manager.py | 229 ++++++++++ .../executor/test_prefix_load_completion.py | 176 ++++++++ 17 files changed, 2760 insertions(+), 189 deletions(-) create mode 100644 tensorrt_llm/_torch/pyexecutor/connectors/prefix_load_completion.py create mode 100644 tests/unittest/_torch/executor/test_kv_connector_executor_lifetime.py create mode 100644 tests/unittest/_torch/executor/test_kv_connector_reservations.py create mode 100644 tests/unittest/_torch/executor/test_prefix_load_completion.py diff --git a/docs/source/features/kv-cache-connector.md b/docs/source/features/kv-cache-connector.md index 902d6844f53a..eca6e78abdd8 100644 --- a/docs/source/features/kv-cache-connector.md +++ b/docs/source/features/kv-cache-connector.md @@ -33,9 +33,18 @@ These methods run on the leader process and drive the connector's behavior. * **Returns**: An arbitrary metadata object (picklable) that describes the tasks for the workers. This object is broadcasted to all workers. * **`get_num_new_matched_tokens(self, request: LlmRequest, num_computed_tokens: int) -> tuple[int, bool]`** - * **Description**: Called when a new request arrives. It checks to see if any KV cache can be loaded from an external KV store. + * **Description**: Queries external KV after the compute batch is selected on the legacy path. Connectors that implement the complete reservation protocol use `reserve_prefix` during admission when prefix-aware scheduling and KV cache manager V2 are enabled. * **Returns**: A tuple `(num_tokens, is_async)`. `num_tokens` is the number of tokens found in the external cache. `is_async` indicates if the loading will happen asynchronously (background) or requires blocking. +* **`reserve_prefix(self, request: LlmRequest, num_computed_tokens: int, reservation_id: int) -> tuple[int, bool]`** + * **Description**: Reserves an additional contiguous prefix beginning at `num_computed_tokens` during batch construction. A positive answer protects that source range against mutation and eviction. The query must start no KV transmission or destination writes. + * **Returns**: `(additional_tokens, is_async)`. The runtime may accept a shorter, block-aligned range and keeps at least the final prompt token for local computation. `is_async=True` parks an accepted request until all workers finish its load. + * **Identity**: The runtime assigns a fresh `reservation_id` to each attempt. Keep reservation state by this identity, including when the same request is retried. + +* **`release_prefix_reservation(self, request: LlmRequest, reservation_id: int, start: int, end: int) -> None`** + * **Description**: Releases the named reservation's protection for the absolute half-open token interval `[start, end)`. The runtime releases rejected or clipped portions before transmission and the accepted portion after all workers report completion. Overlapping reservations must keep their own protection. + * **Lifetime**: This callback never asks the connector to interrupt an active transfer. Client cancellation after dispatch drains the load before releasing its source and destination resources. + * **`request_finished(self, request: LlmRequest, cache_block_ids: list[int]) -> bool`** * **Description**: Called when a request completes generation. * **Returns**: A boolean indicating if an asynchronous save operation is underway. If `True`, the system waits for the operation to complete before releasing the KV cache blocks. @@ -175,13 +184,13 @@ KV pool; the scheduler's exhaustion error says so directly when a connector is a The connector holds page indices across iterations, and `RequestData` reports only the pages appended since the last call. Anything that hands a slot the connector already knows about to a different -request therefore goes unreported and corrupts the next transfer against it. Three configurations do -that, and each is rejected at bring-up. +request therefore requires resetting the connector state before replay. The runtime applies the +following configuration limits at bring-up. | Configuration | Mechanism | |---|---| | Speculative decoding | Rejected draft tokens shrink a request's page list, and the freed slot goes to whichever request allocates next. The connector is never told the tail block moved. | -| A capacity scheduler policy other than `GUARANTEED_NO_EVICT` | A destroyed-and-replayed request comes back on different pages, and the connector's per-request block delta is then measured against pages that were freed with it. | +| A capacity scheduler policy other than `GUARANTEED_NO_EVICT` with KV cache manager V1 | V1 does not reset the connector block delta when a request is destroyed and replayed. V2 resets that state and permits replay; an active connector load retains its allocation until completion. | | A host or disk cache tier | Tier eviction reassigns the GPU slot. See [KV cache tiers are GPU-only under a connector](#kv-cache-tiers-are-gpu-only-under-a-connector). | The exact set the runtime refuses depends on your cache configuration; the bring-up error is @@ -238,16 +247,60 @@ Variable sliding-window attention is the case where that stops working, because A model whose layers all share one sliding window stays a single layer group, so the flat callbacks still apply and such a connector is not refused. What differs is that the callbacks cover the **live window only**. Blocks the window has passed report `-1` (`BAD_PAGE_INDEX`) in place — the list stays aligned to block ordinals, so entry `i` still describes tokens `[i * tokens_per_block, (i+1) * tokens_per_block)`, but the up-front blocks carry no page and are not available to load into or save from. Filter with `valid_page_slots`, described in [Blocks with no page](#blocks-with-no-page); without it a `-1` resolves to the last page slot of the pool. A warning naming the window size is logged at start-up when a flat-only connector is attached to such a model. -`get_num_new_matched_tokens` is asked once the batch for the upcoming forward pass is final. A request that is asked is therefore a request that runs, and the connector can take ownership of remote blocks in the query and release it in `request_finished`. +##### Reserving a prefix during admission + +For connector implementers, scheduler participation requires three optional methods together: +`reserve_prefix` and `release_prefix_reservation` on the scheduler, and +`get_finished_prefix_loads` on the worker. Bring-up rejects a partial implementation. KV cache +manager V2 enables this protocol when `SchedulerConfig.enable_prefix_aware_scheduling=True`. +Connectors using the existing methods retain the final-batch query path. + +The scheduler charges compute tokens after the local and reserved external prefix. The complete +prefix still needs destination KV pages, so a token-budget hit can remain limited by KV capacity. +After the final admission checks, `SchedulerOutput.prefix_loads` authorizes the exact loads: + +| `PrefixLoad` field | Meaning | +|---|---| +| `reservation_id` | The identity supplied to `reserve_prefix`, also returned on completion. | +| `request_id` | The request owning the allocation. | +| `start`, `end` | Absolute half-open token interval to load. | +| `is_async` | Whether the request waits outside the compute batch. | +| `block_ids_by_layer_group` | Complete destination page lists, indexed by group and block ordinal. Filter entries through `valid_page_slots`. | +| `tokens`, `cache_salt` | Request tokens and cache-key isolation salt. | + +Build transfers from this collection. A confirmed async load appears here even when it is absent +from `new_requests` and `cached_requests`, and even when the compute batch is empty. Those two +lists continue to describe computation and its token/block deltas. The reservation itself supplies +no permission to write destination memory. + +| Event | Connector and runtime behavior | +|---|---| +| Candidate queried | Connector protects the promised source; no transmission starts. | +| Candidate rejected or offer shortened | Runtime releases the unused source interval. A later attempt gets a new reservation identity. | +| Accepted load dispatched | Connector starts only the confirmed interval; runtime retains its destination allocation. | +| Client cancels during transmission | Runtime keeps the allocation while the connector completes the load. | +| Every worker completes | Runtime releases the source reservation, resumes a live request, or finalizes a cancelled request and frees its allocation. | + +Each worker reports completion after its destination writes and source reads finish. Reports are +sent to the leader without a collective. Once every worker has reported, an ordered control item +retires the load on all ranks before scheduling. This releases the source reservation and allows +an async request to rejoin the compute batch, or a cancelled request to release its allocation. +The runtime continues polling when no forward pass is scheduled. A transfer that never completes +retains its resources; elapsed time alone cannot make pages safe to reuse. -Two things are worth knowing when tuning a deployment. +For a load admitted with computation, `wait_for_layer_load` establishes the dependency on that +worker's transfer stream. Inference does not wait for the retirement control item. -* **The runtime may honour less than you offer.** With chunked prefill the cache is allocated per context chunk, which is what bounds its memory, so an offer reaching past the current chunk requires the runtime to grow the allocation and that can fail under pressure. The runtime then serves the part it can cover and computes the rest locally. The amount actually served is what `RequestData.computed_position` reflects; the unserved remainder needs no action from the connector beyond its usual `request_finished` cleanup. -* **The query is not part of the scheduler's budget.** The scheduler sizes a request's chunk as if the connector will serve nothing, so a served prefix reduces the work in the forward pass but does not free budget for another request in the same iteration. +##### Legacy final-batch queries -Specify `enable_block_reuse=True` alongside the connector for any of this to run; see [Block reuse alongside the connector](#block-reuse-alongside-the-connector). +`get_num_new_matched_tokens` is called at most once per KV allocation after batch selection. +Its offer can reduce computation for that request, but does not free token budget for another +request in the same iteration. With chunked prefill, allocation growth can limit the served prefix; +`RequestData.computed_position` reflects the load interval that the runtime honors. -`get_num_new_matched_tokens` is called **at most once per KV allocation**. This is the precise form of the "once per request" rule: if a request's KV cache is destroyed and the request is replayed -- which `MAX_UTILIZATION` does under memory pressure -- the replay asks again, because the pages the first answer described are gone. +Destroying an allocation clears its connector state, so replay queries again. Specify +`enable_block_reuse=True` alongside the connector; see +[Block reuse alongside the connector](#block-reuse-alongside-the-connector). **Deployment note.** Under a connector, a workload that was token-bound becomes KV-bound: the connector removes forward-pass tokens but its prefix still occupies GPU pages. Lowering `max_num_tokens` to hand memory back to the KV pool is usually the right adjustment, the opposite of the guidance for a connector-free deployment. @@ -280,6 +333,10 @@ These methods run on all workers (GPU processes) and interact with the actual GP * **Description**: Polled by the runtime to check the status of asynchronous operations. * **Returns**: Two lists of request IDs: those that have finished saving, and those that have finished loading. +* **`get_finished_prefix_loads(self) -> list[int]`** + * **Description**: Reports locally completed reservation identities for loads dispatched through `SchedulerOutput.prefix_loads`, including synchronous loads. Completion means this worker has finished all reads and writes and established the required CUDA stream visibility. This method must not wait for other workers. The runtime collects reports asynchronously and distributes retirement decisions before scheduling; parked async requests resume and shared resources are released only after every worker has reported. + * **Compatibility**: Legacy request-ID load and save completions continue through `get_finished`. + ## Example Implementation The file `examples/llm-api/llm_kv_cache_connector.py` provides a reference implementation of a **Persistent KV Cache**. @@ -287,17 +344,19 @@ The file `examples/llm-api/llm_kv_cache_connector.py` provides a reference imple ### Overview This example implements a file-system based KV cache. -1. **Save**: When a request finishes or needs to be swapped out, its KV blocks are saved to disk as `.pt` files. +1. **Save**: After prefill, complete computed KV blocks are saved to disk as immutable `.pt` files. 2. **Load**: When a new request arrives with the same prompt prefix, the connector identifies the cached files and loads them back into GPU memory, skipping re-computation. ### Implementation Details -* **Metadata**: The example defines a `PersistentKvCacheConnectorMetadata` dataclass containing lists of `(file_path, block_id)` tuples for both loading and saving. This simple structure allows the Scheduler to tell the Worker exactly which file corresponds to which GPU block index. +* **Metadata**: `PersistentKvCacheConnectorMetadata` carries `(file_path, block_id)` load/save targets and accepted reservation identities. The worker returns those identities after the copies finish. + +* **Source protection**: Saves publish complete files atomically and never overwrite an existing inode. `reserve_prefix` creates private hard links to these immutable files without reading KV data. Releasing a range removes only its reservation links, so another reservation keeps its source available. External cache management must preserve this immutability: remove or replace a path rather than writing into a published file. -* **Hashing Strategy**: The `PersistentKvCacheConnectorLeader` hashes the token sequence of a block to generate a unique filename (e.g., `hash_value.pt`). This acts as the lookup key. +* **Hashing Strategy**: `PersistentKvCacheConnectorLeader` hashes the entire token prefix through each block, together with `cache_salt`, using SHA-256. This distinguishes identical blocks reached through different preceding tokens and remains stable across Python processes. Use separate cache directories for different models and KV layouts. * **Worker Logic**: - * `start_load_kv`: Iterates through the load list provided in the metadata, loads the `.pt` file to CPU, and copies it to the specific `block_id` in the GPU tensor. + * `start_load_kv`: Reads only accepted load targets, copies their `.pt` data to GPU with blocking copies, and records their reservation identities for `get_finished_prefix_loads`. * `wait_for_save`: Performs the reverse. It copies data from the GPU `block_id` to CPU and saves it to disk using `torch.save`. ### Limitations & Patterns @@ -305,7 +364,7 @@ This example implements a file-system based KV cache. This example illustrates the API mechanics but has several limitations that make it unsuitable for high-performance production use without modification: 1. **Blocking I/O**: The example uses `torch.load` and `torch.save` synchronously. In a real implementation, these should be offloaded to a background thread or asynchronous I/O handler to avoid stalling the GPU. -2. **Simplified Block Matching**: The `get_num_new_matched_tokens` implementation in the example only matches full blocks. It does not handle partial cache hits. +2. **Simplified Block Matching**: The example matches and loads complete blocks, requires one layer group, and saves only the first context chunk. It does not demonstrate chunked-prefill persistence. 3. **FileSystem Latency**: Storing one file per block can create high filesystem overhead. ### Usage diff --git a/examples/llm-api/llm_kv_cache_connector.py b/examples/llm-api/llm_kv_cache_connector.py index 73e847fa3425..66cbb8fd08e6 100644 --- a/examples/llm-api/llm_kv_cache_connector.py +++ b/examples/llm-api/llm_kv_cache_connector.py @@ -1,3 +1,18 @@ +# SPDX-FileCopyrightText: Copyright (c) 2022-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. + ### :title KV Cache Connector ### :order 6 ### :section Customization @@ -74,17 +89,18 @@ - Cache files are stored in a temporary directory (cleaned up after the demo) - The implementation is simplified and not optimized for production use - Does not support chunked prefill in this example -- See `tensorrt_llm/_torch/pyexecutor/kv_cache_connector.py` for the full connector interface +- See `tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_connector.py` for the full connector interface **NOTE:** This example connector implementation is designed for demonstration purposes and is NOT suitable for production use without additional optimizations and error handling. ''' +import hashlib import os import sys from dataclasses import dataclass, field from pathlib import Path -from tempfile import TemporaryDirectory +from tempfile import NamedTemporaryFile, TemporaryDirectory from typing import Optional import click @@ -105,6 +121,7 @@ class PersistentKvCacheConnectorMetadata: load: list[tuple[str, int]] = field(default_factory=list) save: list[tuple[str, int]] = field(default_factory=list) + prefix_load_ids: list[int] = field(default_factory=list) class PersistentKvCacheConnectorWorker(KvCacheConnectorWorker): @@ -113,6 +130,7 @@ def __init__(self, llm_args: TorchLlmArgs): super().__init__(llm_args) self.kv_cache_tensor = None + self._finished_prefix_loads: list[int] = [] def register_kv_caches(self, kv_cache_tensor: torch.Tensor): # This is the only registration hook this connector needs. A cache that @@ -131,6 +149,12 @@ def start_load_kv(self, stream: torch.cuda.Stream): # Copy into the device block. self.kv_cache_tensor[block_id].copy_(cpu_tensor, non_blocking=False) + self._finished_prefix_loads.extend(self._metadata.prefix_load_ids) + + def get_finished_prefix_loads(self) -> list[int]: + finished, self._finished_prefix_loads = self._finished_prefix_loads, [] + return finished + def wait_for_layer_load(self, layer_idx: int, stream: torch.cuda.Stream): pass @@ -149,8 +173,14 @@ def wait_for_save(self, stream: torch.cuda.Stream): if Path(path).exists(): continue - # Do a blocking save to the file. This way, we only return once all saves are complete. - torch.save(cpu_tensor, path) + # Publish complete immutable files so a reservation can retain an + # inode while another process replaces or removes its cache key. + with NamedTemporaryFile(dir=Path(path).parent) as staging: + torch.save(cpu_tensor, staging.name) + try: + os.link(staging.name, path) + except FileExistsError: + pass def get_finished( self, finished_gen_req_ids: list[int], @@ -171,70 +201,68 @@ def __init__(self, llm_args: TorchLlmArgs): "./connector_cache") os.makedirs(self.cache_folder, exist_ok=True) + self._reservation_folder = TemporaryDirectory(prefix=".reservations-", + dir=self.cache_folder) + self._reserved_files: dict[int, dict[int, Path]] = {} + self._reserved_ranges: dict[int, list[tuple[int, int]]] = {} def build_connector_meta(self, scheduler_output: SchedulerOutput): # NOTE: This is a simplified implementation, and does not work with chunked prefill. metadata = PersistentKvCacheConnectorMetadata() - for req in scheduler_output.new_requests: - # If we don't have any pending loads for this request, we can skip it. - if req.request_id not in self.pending_loads: - continue - - num_computed_blocks = req.computed_position // self.block_size - block_ids = req.new_block_ids - - pending_load = self.pending_loads[req.request_id] - - # Ordinal -> page slot for the blocks that have a page. Blocks with - # none keep their ordinal in `block_ids` so that entry `i` always - # describes the same token range; they are dropped here so no - # transfer can be built against one. + accepted_ends = {} + for load in scheduler_output.prefix_loads: + if len(load.block_ids_by_layer_group) != 1: + raise ValueError( + "Persistent connector requires one layer group") + if load.start % self.block_size or load.end % self.block_size: + raise ValueError("Persistent connector loads complete blocks") + block_ids = load.block_ids_by_layer_group[0] + if load.end // self.block_size > len(block_ids): + raise ValueError( + "Confirmed prefix exceeds destination allocation") slots = dict(valid_page_slots(block_ids)) + files = self._reserved_files[load.reservation_id] + for ordinal in range(load.start // self.block_size, + load.end // self.block_size): + if ordinal in slots: + metadata.load.append((str(files[ordinal]), slots[ordinal])) + metadata.prefix_load_ids.append(load.reservation_id) + accepted_ends[load.request_id] = load.end - for file_path, block_pos in zip( - pending_load, range(num_computed_blocks, len(block_ids))): - slot = slots.get(block_pos) - if slot is None: - continue - metadata.load.append((file_path, slot)) - - # Break up the remainder of the token sequence into chunks. - chunks = self._chunk_tokens(req.new_tokens) - - # For each chunk that isn't already on device, and isn't in our connector cache, we need to save it. - for block_pos in range(num_computed_blocks + len(pending_load), - len(block_ids)): - slot = slots.get(block_pos) - if slot is None: + for req in scheduler_output.new_requests: + num_computed_blocks = req.computed_position // self.block_size + slots = dict(valid_page_slots(req.new_block_ids)) + pending_load = self.pending_loads.get(req.request_id, []) + for ordinal, path in enumerate(pending_load, num_computed_blocks): + if ordinal in slots: + metadata.load.append((str(path), slots[ordinal])) + + loaded_end = accepted_ends.get( + req.request_id, + (num_computed_blocks + len(pending_load)) * self.block_size) + computed_end = max(req.computed_position, + loaded_end) + req.num_scheduled_tokens + for ordinal, slot in slots.items(): + end = (ordinal + 1) * self.block_size + if end <= loaded_end or end > min(len(req.new_tokens), + computed_end): continue - if len(chunks[block_pos]) == self.block_size: - hashed_tokens = self._hash_tokens(chunks[block_pos], - req.cache_salt) - - file_path = self._file_path(hashed_tokens) - - metadata.save.append((file_path, slot)) + key = self._hash_tokens(req.new_tokens[:end], req.cache_salt) + metadata.save.append((str(self._file_path(key)), slot)) self.pending_loads = {} - return metadata - def _hash_tokens(self, tokens: list[int], cache_salt: Optional[str]) -> int: - # cache_salt must participate in the hash so that requests carrying - # different salts (or no salt) cannot collide on the same cache file. - return abs(hash((cache_salt, tuple(tokens)))) + def _hash_tokens(self, tokens: list[int], cache_salt: Optional[str]) -> str: + # KV depends on every preceding token, including those in other blocks. + key = repr((cache_salt, tuple(tokens))).encode("utf-8") + return hashlib.sha256(key).hexdigest() - def _file_path(self, hash_value: int) -> Path: + def _file_path(self, hash_value: str) -> Path: return Path(self.cache_folder) / f"{hash_value}.pt" - def _chunk_tokens(self, tokens: list[int]) -> list[list[int]]: - return [ - tokens[i:i + self.block_size] - for i in range(0, len(tokens), self.block_size) - ] - def get_num_new_matched_tokens( self, request: LlmRequest, num_computed_tokens: int) -> tuple[int, bool]: @@ -244,28 +272,14 @@ def get_num_new_matched_tokens( if (num_computed_tokens % self.block_size) != 0: return 0, False - computed_blocks = num_computed_tokens // self.block_size - - # Get all the tokens that don't have a cache hit on device. - remaining_tokens = request.get_tokens(0)[computed_blocks * - self.block_size:] - - remaining_chunks = self._chunk_tokens(remaining_tokens) - - # For each chunk, check if it exists in our cache. - for chunk in remaining_chunks: - # Only do full blocks. - if len(chunk) == self.block_size: - hashed_tokens = self._hash_tokens(chunk, request.cache_salt) - - file_path = self._file_path(hashed_tokens) - - # If we get a cache hit, we want to load it into device. - # Otherwise, we can stop looking. - if file_path.exists(): - self.pending_loads[request.request_id].append(file_path) - else: - break + tokens = request.get_tokens(0) + for end in range(num_computed_tokens + self.block_size, + len(tokens) + 1, self.block_size): + key = self._hash_tokens(tokens[:end], request.cache_salt) + file_path = self._file_path(key) + if not file_path.exists(): + break + self.pending_loads[request.request_id].append(file_path) logger.info( f"KV CONNECTOR: Matched {len(self.pending_loads[request.request_id])} blocks for request {request.request_id}" @@ -274,6 +288,63 @@ def get_num_new_matched_tokens( return len( self.pending_loads[request.request_id]) * self.block_size, False + def reserve_prefix(self, request: LlmRequest, num_computed_tokens: int, + reservation_id: int) -> tuple[int, bool]: + """Protect immutable source files without reading their KV data.""" + if num_computed_tokens % self.block_size: + return 0, False + tokens = request.get_tokens(0) + folder = Path(self._reservation_folder.name) / str(reservation_id) + folder.mkdir() + files = {} + for end in range(num_computed_tokens + self.block_size, + len(tokens) + 1, self.block_size): + ordinal = end // self.block_size - 1 + key = self._hash_tokens(tokens[:end], request.cache_salt) + protected = folder / f"{ordinal}.pt" + try: + os.link(self._file_path(key), protected) + except FileNotFoundError: + break + files[ordinal] = protected + count = len(files) * self.block_size + if count: + self._reserved_files[reservation_id] = files + self._reserved_ranges[reservation_id] = [ + (num_computed_tokens, num_computed_tokens + count) + ] + else: + folder.rmdir() + return count, False + + def release_prefix_reservation(self, request: LlmRequest, + reservation_id: int, start: int, + end: int) -> None: + """Release only the named range; overlapping reservations retain it.""" + remaining = [] + for left, right in self._reserved_ranges[reservation_id]: + if right <= start or left >= end: + remaining.append((left, right)) + else: + if left < start: + remaining.append((left, start)) + if end < right: + remaining.append((end, right)) + files = self._reserved_files[reservation_id] + for ordinal, path in list(files.items()): + block_start = ordinal * self.block_size + block_end = block_start + self.block_size + if not any(left < block_end and right > block_start + for left, right in remaining): + path.unlink() + del files[ordinal] + if remaining: + self._reserved_ranges[reservation_id] = remaining + else: + del self._reserved_ranges[reservation_id] + del self._reserved_files[reservation_id] + (Path(self._reservation_folder.name) / str(reservation_id)).rmdir() + def request_finished(self, request: LlmRequest, cache_block_ids: list[int]) -> bool: # We don't do any asynchronous saving, so always return False diff --git a/tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_connector.py b/tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_connector.py index 99bdff750e77..73e88a52ef0c 100644 --- a/tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_connector.py +++ b/tensorrt_llm/_torch/pyexecutor/connectors/kv_cache_connector.py @@ -41,7 +41,7 @@ import torch -from tensorrt_llm._utils import mpi_allgather, mpi_broadcast, mpi_rank +from tensorrt_llm._utils import mpi_allgather, mpi_broadcast, mpi_comm, mpi_rank, mpi_world_size from tensorrt_llm.bindings import LlmRequestState from tensorrt_llm.bindings.internal.batch_manager import ( KvCacheConnectorManager as KvCacheConnectorManagerCpp, @@ -52,6 +52,7 @@ from ..llm_request import get_draft_token_length from ..scheduler import ScheduledRequests +from .prefix_load_completion import PrefixLoadCompletionTracker if TYPE_CHECKING: from ..resource_manager import KVCacheManager @@ -112,6 +113,26 @@ class SchedulerOutput: # Requests being scheduled, that have already shown up in `new_requests`. cached_requests: List[RequestData] = field(default_factory=list) + # Confirmed transfers, including loads whose requests are parked outside + # the compute batch. A reservation ID identifies one destination allocation. + prefix_loads: List["PrefixLoad"] = field(default_factory=list) + + +@dataclass(frozen=True) +class PrefixReservation: + reservation_id: int + request_id: int + start: int + end: int + is_async: bool + + +@dataclass(frozen=True) +class PrefixLoad(PrefixReservation): + block_ids_by_layer_group: List[List[int]] + tokens: List[int] + cache_salt: Optional[str] + def _flat_form_unavailable(flat: str, grouped: str) -> Callable: """A stand-in for ``flat`` that names the form this connector implements.""" @@ -292,6 +313,16 @@ def get_finished( longer than others to complete the operations. """ + def get_finished_prefix_loads(self) -> List[int]: + """Return reservation IDs whose confirmed writes have finished locally. + + Implement together with the scheduler's reservation methods. Report + both synchronous and asynchronous loads, using the IDs received through + ``SchedulerOutput.prefix_loads``. Do not wait for other workers. The + runtime retains allocations until an ordered retirement decision arrives. + """ + return [] + class KvCacheConnectorScheduler(ABC): def __init__(self, llm_args: TorchLlmArgs): @@ -338,6 +369,29 @@ def get_num_new_matched_tokens( Whether the tokens will be loaded asynchronously. """ + def reserve_prefix( + self, request: LlmRequest, num_computed_tokens: int, reservation_id: int + ) -> Tuple[int, bool]: + """Protect an additional prefix without starting a transfer. + + The returned count defines ``[num_computed_tokens, start + count)``. + Keep that content immutable and available until its exact range is + released. Loads start only from confirmed ``prefix_loads`` metadata; + the runtime may release all or part of this promise before admission. + """ + raise NotImplementedError + + def release_prefix_reservation( + self, request: LlmRequest, reservation_id: int, start: int, end: int + ) -> None: + """Release protection for a half-open range of a reservation. + + The range is either unused or complete on every worker. This callback + never aborts transmission, and overlapping reservations retain their + own protection until each is released. + """ + raise NotImplementedError + @abstractmethod def request_finished(self, request: LlmRequest, cache_block_ids: List[int]) -> bool: """ @@ -673,6 +727,20 @@ def __init__( self._scheduler_output = None self.scheduler_output_manager = KvCacheConnectorSchedulerOutputManager() + self.prefix_reservations_enabled = False + self._prefix_capability = None + self._next_prefix_reservation_id = 1 + self._prefix_reservations: Dict[int, PrefixReservation] = {} + self._prefix_requests: Dict[int, LlmRequest] = {} + self._prefix_loads: Dict[int, PrefixLoad] = {} + self._prefix_load_requests: Dict[int, LlmRequest] = {} + self._bound_scheduler_output: Optional[SchedulerOutput] = None + self._bound_prefix_load_ids: Set[int] = set() + self._dispatched_prefix_load_ids: Set[int] = set() + self._local_finished_prefix_load_ids: Set[int] = set() + self._prefix_completion_tracker: PrefixLoadCompletionTracker | None = None + self._deferred_load_terminations: Dict[int, LlmRequest] = {} + self._finished_load_terminations: List[LlmRequest] = [] def _run_on_leader(self, f: Callable[[], Any]) -> Any: """ @@ -685,10 +753,224 @@ def _run_on_leader(self, f: Callable[[], Any]) -> Any: res = None return mpi_broadcast(res, root=0) + def configure_prefix_reservations(self, enabled: bool) -> None: + """Enable reservations when both connector roles implement the protocol.""" + if self._prefix_capability is None: + scheduler_methods = self._run_on_leader( + lambda: [ + getattr(type(self.scheduler), name, None) is not None + and getattr(type(self.scheduler), name) + is not getattr(KvCacheConnectorScheduler, name) + for name in ("reserve_prefix", "release_prefix_reservation") + ] + ) + worker_method = getattr(type(self.worker), "get_finished_prefix_loads", None) + worker_capabilities = mpi_allgather( + worker_method is not None + and worker_method is not KvCacheConnectorWorker.get_finished_prefix_loads + ) + capabilities = scheduler_methods + worker_capabilities + if any(capabilities) and not all(capabilities): + raise ValueError( + "KV connector prefix reservations require reserve_prefix and " + "release_prefix_reservation on the scheduler and " + "get_finished_prefix_loads on every worker." + ) + self._prefix_capability = all(capabilities) + self.prefix_reservations_enabled = enabled and self._prefix_capability + if self.prefix_reservations_enabled and self._prefix_completion_tracker is None: + comm = mpi_comm().Dup() if mpi_world_size() > 1 else None + self._prefix_completion_tracker = PrefixLoadCompletionTracker(comm) + + def reserve_prefix(self, request: LlmRequest, local_end: int) -> Optional[PrefixReservation]: + """Query once per pending attempt and protect the promised source range.""" + if not self.prefix_reservations_enabled: + return None + if local_end < 0: + raise ValueError("A prefix reservation cannot start before the prompt") + existing = self.get_prefix_reservation(request) + if existing is not None: + return existing + if self.has_pending_load(request): + raise RuntimeError("Cannot reserve a prefix while a load owns this allocation") + reservation_id = self._next_prefix_reservation_id + self._next_prefix_reservation_id += 1 + count, is_async = self._run_on_leader( + lambda: self.scheduler.reserve_prefix(request, local_end, reservation_id) + ) + if not isinstance(count, int) or isinstance(count, bool) or count < 0: + raise ValueError("A prefix reservation must return a nonnegative token count") + if not isinstance(is_async, bool): + raise ValueError("A prefix reservation must return a boolean asynchronous mode") + if count == 0: + if is_async: + raise ValueError("An empty prefix reservation cannot load asynchronously") + return None + reservation = PrefixReservation( + reservation_id, request.request_id, local_end, local_end + count, is_async + ) + self._prefix_reservations[request.request_id] = reservation + self._prefix_requests[request.request_id] = request + logger.debug( + f"KV connector reserved {reservation_id} for request {request.request_id}: " + f"[{reservation.start}, {reservation.end}), async={is_async}" + ) + return reservation + + def get_prefix_reservation(self, request: LlmRequest) -> Optional[PrefixReservation]: + return self._prefix_reservations.get(request.request_id) + + def _release_prefix_range( + self, request: LlmRequest, reservation: PrefixReservation, start: int, end: int + ) -> None: + if start < end and self.scheduler is not None: + self.scheduler.release_prefix_reservation( + request, reservation.reservation_id, start, end + ) + + def trim_prefix_reservation( + self, request: LlmRequest, start: int, end: int + ) -> Optional[PrefixReservation]: + """Keep only the requested subrange and release both discarded ends.""" + reservation = self.get_prefix_reservation(request) + if reservation is None: + return None + if not reservation.start <= start <= end <= reservation.end: + raise ValueError("The accepted prefix must be inside its reservation") + self._release_prefix_range(request, reservation, reservation.start, start) + self._release_prefix_range(request, reservation, end, reservation.end) + if start == end: + self._prefix_reservations.pop(request.request_id) + self._prefix_requests.pop(request.request_id) + return None + trimmed = PrefixReservation( + reservation.reservation_id, reservation.request_id, start, end, reservation.is_async + ) + self._prefix_reservations[request.request_id] = trimmed + return trimmed + + def release_prefix_reservation(self, request: LlmRequest) -> None: + """Release an unaccepted promise without affecting a confirmed transfer.""" + reservation = self._prefix_reservations.get(request.request_id) + if reservation is not None: + self._release_prefix_range(request, reservation, reservation.start, reservation.end) + self._prefix_reservations.pop(request.request_id) + self._prefix_requests.pop(request.request_id) + + def pending_prefix_requests(self) -> List[LlmRequest]: + return list(self._prefix_requests.values()) + + def accept_prefix_load( + self, + request: LlmRequest, + start: int, + end: int, + block_ids_by_layer_group: List[List[int]], + ) -> None: + """Confirm a transfer against an allocation that survives until completion.""" + reservation = self.trim_prefix_reservation(request, start, end) + if reservation is None: + return + load = PrefixLoad( + reservation.reservation_id, + reservation.request_id, + reservation.start, + reservation.end, + reservation.is_async, + [list(indices) for indices in block_ids_by_layer_group], + list(request.get_tokens(0)), + request.cache_salt, + ) + self._prefix_loads[load.reservation_id] = load + self._prefix_load_requests[request.request_id] = request + self._prefix_completion_tracker.track(load.reservation_id) + logger.debug( + f"KV connector accepted load {load.reservation_id} for request {request.request_id}: " + f"[{start}, {end})" + ) + self._prefix_reservations.pop(request.request_id) + self._prefix_requests.pop(request.request_id) + if not load.is_async: + self.commit_new_matched_tokens(request, end - start, False) + else: + request.py_num_connector_matched_tokens = end - start + + def mark_prefix_loads_dispatched(self) -> None: + """Retain destination ownership before the worker can enqueue a write.""" + self._dispatched_prefix_load_ids.update(self._bound_prefix_load_ids) + self._bound_prefix_load_ids.clear() + self._bound_scheduler_output = None + + def has_pending_load(self, request: LlmRequest) -> bool: + request_id = request.request_id + return ( + request_id in self._prefix_load_requests + or request_id in self.new_async_requests.loading + or request_id in self.pending_async_requests.loading + or request_id in self.local_finished_async_requests.loading + ) + + def has_pending_loads(self) -> bool: + return bool( + self._prefix_load_requests + or self._finished_load_terminations + or self.new_async_requests.loading + or self.pending_async_requests.loading + or self.local_finished_async_requests.loading + ) + + def release_unstarted_prefix_loads(self, request: LlmRequest) -> None: + """Abandon accepted work only while no worker has been authorized to write.""" + for reservation_id, load in list(self._prefix_loads.items()): + if ( + load.request_id != request.request_id + or reservation_id in self._dispatched_prefix_load_ids + ): + continue + self._release_prefix_range(request, load, load.start, load.end) + del self._prefix_loads[reservation_id] + self._prefix_completion_tracker.forget(reservation_id) + self._prefix_load_requests.pop(request.request_id) + self._bound_prefix_load_ids.discard(reservation_id) + self.scheduler_output_manager.external_loads.pop(request.request_id, None) + for output in (self._scheduler_output, self._bound_scheduler_output): + if output is None: + continue + output.prefix_loads = [ + item for item in output.prefix_loads if item.reservation_id != reservation_id + ] + output.new_requests = [ + item for item in output.new_requests if item.request_id != request.request_id + ] + output.cached_requests = [ + item for item in output.cached_requests if item.request_id != request.request_id + ] + if self._bound_scheduler_output is not None: + metadata = self._run_on_leader( + lambda: self.scheduler.build_connector_meta(self._bound_scheduler_output) + ) + self.worker.bind_connector_meta(metadata) + + def defer_load_termination(self, request: LlmRequest) -> bool: + """Remember termination while a load still owns the request's pages.""" + if not self.has_pending_load(request): + return False + if request.request_id not in self._deferred_load_terminations: + logger.debug( + f"KV connector draining load before terminating request {request.request_id}" + ) + self._deferred_load_terminations[request.request_id] = request + return True + + def take_finished_load_terminations(self) -> List[LlmRequest]: + finished = self._finished_load_terminations + self._finished_load_terminations = [] + return finished + def query_num_new_matched_tokens( self, request: LlmRequest, num_computed_tokens: int ) -> Tuple[int, bool]: - """Ask the connector how much of the prompt it can serve. No side effects. + """Query the legacy connector without committing runtime load bookkeeping. The connector ABC promises one query per allocation, so a caller must reach this at most once per request per allocation whatever it does with @@ -847,15 +1129,13 @@ def warn_flat_scheduler_under_swa(self, window_size: Optional[int]) -> None: ) def reset_request_state(self, request: LlmRequest) -> None: - """Tell the connector bookkeeping that this request's allocation died. - - Only a cache that can destructively pause and replay a live request - needs this; under ``GUARANTEED_NO_EVICT`` an allocation dies only when - the request finishes. ``KVCacheV2Scheduler`` coerces the policy to - ``MAX_UTILIZATION`` whatever was configured, so everything keyed to the - old allocation has to go with it. - """ + """Retire bookkeeping only after the allocation's writes have completed.""" + if self.has_pending_load(request): + raise RuntimeError("Cannot reset connector state while a load owns the allocation") + self.release_prefix_reservation(request) self.scheduler_output_manager.reset_request(request.request_id) + self.finished_async_loading_requests.pop(request.request_id, None) + self._deferred_load_terminations.pop(request.request_id, None) def should_add_sequence(self, request: LlmRequest) -> bool: req_id = request.request_id @@ -864,14 +1144,30 @@ def should_add_sequence(self, request: LlmRequest) -> bool: def build_scheduler_output( self, scheduled_batch: ScheduledRequests, kv_cache_manager: "KVCacheManager" ): + async_requests = AsyncRequests( + {}, + { + **self.new_async_requests.loading, + **{ + load.request_id: self._prefix_load_requests[load.request_id] + for load in self._prefix_loads.values() + if load.is_async + }, + }, + ) self._scheduler_output = self.scheduler_output_manager.build_scheduler_output( - scheduled_batch, self.new_async_requests, kv_cache_manager + scheduled_batch, async_requests, kv_cache_manager ) + self._scheduler_output.prefix_loads = [ + load + for reservation_id, load in self._prefix_loads.items() + if reservation_id not in self._dispatched_prefix_load_ids + ] def take_scheduled_requests_pending_load(self, scheduled_requests: ScheduledRequests): """ Remove context requests from our list of scheduled requests that are being loaded asynchronously. - This is done to prevent the runtime from attempting to load the KV cache for these requests. + Their destination pages remain owned while computation waits for completion. Args: scheduled_requests: The scheduled requests. @@ -880,17 +1176,24 @@ def take_scheduled_requests_pending_load(self, scheduled_requests: ScheduledRequ The scheduled requests with the context requests that are being loaded asynchronously removed. """ + prefix_loading_ids = { + load.request_id for load in self._prefix_loads.values() if load.is_async + } for key in ["context_requests_chunking", "context_requests_last_chunk"]: allowed_context_requests = [] for req in getattr(scheduled_requests, key): # If this request is being loaded asynchronously, in # addition to removing it from the list of scheduled # requests, we also need to update its state. - if req.request_id in self.new_async_requests.loading.keys(): + prefix_loading = req.request_id in prefix_loading_ids + if req.request_id in self.new_async_requests.loading or prefix_loading: req.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS # Replace the request with the canonical request. - self.new_async_requests.loading[req.request_id] = req + if prefix_loading: + self._prefix_load_requests[req.request_id] = req + else: + self.new_async_requests.loading[req.request_id] = req else: allowed_context_requests.append(req) setattr(scheduled_requests, key, allowed_context_requests) @@ -903,6 +1206,10 @@ def handle_metadata(self) -> object: lambda: self.scheduler.build_connector_meta(self._scheduler_output) ) + self._bound_prefix_load_ids.update( + load.reservation_id for load in self._scheduler_output.prefix_loads + ) + self._bound_scheduler_output = self._scheduler_output self._scheduler_output = None self.worker.bind_connector_meta(metadata) @@ -977,7 +1284,8 @@ def get_finished(self) -> List[LlmRequest]: # Remove the requests from our pending list that have finished locally. new_local_finished_async_requests = self.pending_async_requests.extract_by_id( - finished_saving, finished_loading + set(finished_saving) & self.pending_async_requests.saving_ids, + set(finished_loading) & self.pending_async_requests.loading_ids, ) # Add these requests to our list of locally finished requests. @@ -999,14 +1307,57 @@ def get_finished(self) -> List[LlmRequest]: ) # For requests that have finished loading, move them back to the context state. - for id, req in all_finished.loading.items(): - req.state = LlmRequestState.CONTEXT_INIT - self.finished_async_loading_requests[id] = req + for req in all_finished.loading.values(): + self._finish_load(req, is_async=True) + + if self.prefix_reservations_enabled: + finished = ( + set(self.worker.get_finished_prefix_loads()) + & self._dispatched_prefix_load_ids - self._local_finished_prefix_load_ids + ) + self._local_finished_prefix_load_ids.update(finished) + self._prefix_completion_tracker.report(finished) + self._prefix_completion_tracker.poll() # Return the requests that have finished saving. # The execution loop will call _terminate_request on these requests. return list(all_finished.saving.values()) + def take_finished_prefix_loads(self) -> list[int]: + """Return leader completion decisions for the next request broadcast.""" + if not self.prefix_reservations_enabled or self.scheduler is None: + return [] + return self._prefix_completion_tracker.take_completed() + + def finish_prefix_loads(self, reservation_ids: list[int]) -> None: + """Apply ordered completion decisions before scheduling on every worker.""" + for reservation_id in reservation_ids: + load = self._prefix_loads.get(reservation_id) + if load is None: + continue + if reservation_id not in self._local_finished_prefix_load_ids: + raise RuntimeError(f"Prefix load {reservation_id} has not finished locally") + request = self._prefix_load_requests[load.request_id] + self._release_prefix_range(request, load, load.start, load.end) + del self._prefix_loads[reservation_id] + del self._prefix_load_requests[load.request_id] + self._dispatched_prefix_load_ids.remove(reservation_id) + self._local_finished_prefix_load_ids.remove(reservation_id) + self._prefix_completion_tracker.forget(reservation_id) + self._finish_load(request, is_async=load.is_async) + + def shutdown(self) -> None: + if self._prefix_completion_tracker is not None: + self._prefix_completion_tracker.close() + + def _finish_load(self, request: LlmRequest, is_async: bool) -> None: + if request.request_id in self._deferred_load_terminations: + request = self._deferred_load_terminations.pop(request.request_id) + self._finished_load_terminations.append(request) + elif is_async: + request.state = LlmRequestState.CONTEXT_INIT + self.finished_async_loading_requests[request.request_id] = request + def update_state_after_alloc( self, req: LlmRequest, @@ -1027,6 +1378,10 @@ def update_state_after_alloc( self.scheduler.update_state_after_alloc(req, block_ids) def set_scheduler_output(self, scheduler_output: SchedulerOutput): + if self._scheduler_output is not None: + loads = {load.reservation_id: load for load in self._scheduler_output.prefix_loads} + loads.update({load.reservation_id: load for load in scheduler_output.prefix_loads}) + scheduler_output.prefix_loads = list(loads.values()) self._scheduler_output = scheduler_output def layer_pre_hook(self, module, *args): diff --git a/tensorrt_llm/_torch/pyexecutor/connectors/prefix_load_completion.py b/tensorrt_llm/_torch/pyexecutor/connectors/prefix_load_completion.py new file mode 100644 index 000000000000..6df87c446669 --- /dev/null +++ b/tensorrt_llm/_torch/pyexecutor/connectors/prefix_load_completion.py @@ -0,0 +1,101 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from mpi4py.util import pkl5 + + +class PrefixLoadCompletionTracker: + """Collect worker reports without waiting for transfers or peer progress. + + The leader publishes completed IDs through the executor's ordered request + queue. Tracking ends when that decision has been applied on every rank. + """ + + def __init__(self, comm: pkl5.Intracomm | None = None) -> None: + self._comm = comm + self._rank = comm.Get_rank() if comm is not None else 0 + self._size = comm.Get_size() if comm is not None else 1 + self._pending: set[int] = set() + self._send_request: pkl5.Request | None = None + self._receive_requests: dict[int, pkl5.Request] = {} + self._workers_finished: dict[int, set[int]] = {} + self._completed: set[int] = set() + + def track(self, reservation_id: int) -> None: + if self._rank == 0: + self._workers_finished[reservation_id] = set() + + def forget(self, reservation_id: int) -> None: + self._workers_finished.pop(reservation_id, None) + self._completed.discard(reservation_id) + self._pending.discard(reservation_id) + + def report(self, reservation_ids: set[int]) -> None: + """Queue newly completed local transfers; sending never waits for the leader.""" + if self._rank == 0: + self._record(0, reservation_ids) + else: + self._pending.update(reservation_ids) + + def _record(self, rank: int, reservation_ids: set[int]) -> None: + for reservation_id in reservation_ids: + workers = self._workers_finished.get(reservation_id) + if workers is None: + continue + workers.add(rank) + if len(workers) == self._size: + self._completed.add(reservation_id) + + def poll(self) -> None: + """Progress at most one report per peer without a blocking receive.""" + if self._comm is None: + return + if self._rank != 0: + if self._send_request is not None: + finished, _ = self._send_request.test() + if not finished: + return + self._send_request = None + if self._pending: + self._send_request = self._comm.isend(self._pending, dest=0, tag=0) + self._pending = set() + return + + if not self._workers_finished and not self._receive_requests: + return + for rank in range(1, self._size): + request = self._receive_requests.get(rank) + if request is None: + message = self._comm.improbe(source=rank, tag=0) + if message is None: + continue + request = message.irecv() + self._receive_requests[rank] = request + finished, reservation_ids = request.test() + if finished: + del self._receive_requests[rank] + self._record(rank, reservation_ids) + + def take_completed(self) -> list[int]: + """Return leader decisions to distribute before the next scheduling pass.""" + self.poll() + completed = sorted(self._completed) + self._completed.clear() + return completed + + def close(self) -> None: + """Release the private communicator after the executor has stopped.""" + if self._send_request is not None: + self._send_request.wait() + self._send_request = None + for request in self._receive_requests.values(): + request.wait() + self._receive_requests.clear() + if self._comm is not None: + self._comm.Free() + self._comm = None diff --git a/tensorrt_llm/_torch/pyexecutor/executor_request_queue.py b/tensorrt_llm/_torch/pyexecutor/executor_request_queue.py index 510fe23fc9a8..f23c82029bdb 100644 --- a/tensorrt_llm/_torch/pyexecutor/executor_request_queue.py +++ b/tensorrt_llm/_torch/pyexecutor/executor_request_queue.py @@ -1,3 +1,6 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + import dataclasses import datetime import enum @@ -19,6 +22,7 @@ # profile window applies on every PyExecutor, not just the leader. PROFILE_START_REQUEST_ID = -3 PROFILE_STOP_REQUEST_ID = -4 +PREFIX_LOAD_COMPLETION_REQUEST_ID = -5 class RequestAdmissionState(enum.Enum): @@ -49,6 +53,7 @@ class RequestQueueItem: # ``num_steps``, ``start_step``, and ``activities`` that every rank # applies when the broadcast reaches them. profile_config: Optional[dict] = None + finished_prefix_load_ids: Optional[list[int]] = None @property def is_shutdown_request(self): @@ -58,7 +63,12 @@ def is_shutdown_request(self): def is_normal_request(self): return not (self.is_shutdown_request or self.is_canceled_request or self.is_control_request or self.is_profile_start_request - or self.is_profile_stop_request) + or self.is_profile_stop_request + or self.is_prefix_load_completion_request) + + @property + def is_prefix_load_completion_request(self): + return self.id == PREFIX_LOAD_COMPLETION_REQUEST_ID @property def is_control_request(self): 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 8a8907d88dc6..eb40fae11e44 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 @@ -3386,6 +3386,11 @@ def revert_allocate_context(self, req: LlmRequest) -> bool: """Undo this iteration's context resize. False means the cache was dropped, not shrunk (history outran pre-resize capacity); the caller drops any draft pool. """ + if self._connector_reservations_enabled(): + if self.kv_connector_manager.get_prefix_reservation(req) is not None: + self.free_resources(req) + rewind_context_after_cache_drop(req, self.tokens_per_block) + return False pre_cap = getattr(req, "py_ctx_pre_resize_cap", None) if pre_cap is None: return True @@ -3544,11 +3549,8 @@ def prepare_context_cache(self, req: LlmRequest, reuse_limit: int | None = None) # scratch blocks are only valid for local prefill chunks. kv_cache.enable_swa_scratch_reuse = False elif self._connector_may_serve(req): - # Same reason, one step earlier: a connector writes real cache - # content into these blocks. Whether it will is not known until - # `prepare_resources` asks, and the flag has to be off before - # `resize_context` can take scratch slots, so it is cleared for - # every servable request rather than only the served ones. + # Connector loads need persistent destination pages; scratch + # slots are reserved for local prefill. kv_cache.enable_swa_scratch_reuse = False if not self._resume_and_restore(req.py_request_id, kv_cache): return None @@ -3565,7 +3567,7 @@ def prepare_context_cache(self, req: LlmRequest, reuse_limit: int | None = None) return kv_cache.num_committed_tokens def prepare_context(self, req: LlmRequest) -> bool: - """Create _KVCache, handle block reuse, and resume. Does NOT resize.""" + """Create/resume the cache and expose local or reserved reuse before budgeting.""" assert not req.is_disagg_generation_init_state, ( f"req {req.py_request_id}: use prepare_disagg_gen_init" ) @@ -3576,6 +3578,7 @@ def prepare_context(self, req: LlmRequest) -> bool: # until context end, so reapplying later would rewind the cursor. if req.is_first_context_chunk and self.enable_block_reuse: _settle_context_cursor(req, reused, self.tokens_per_block) + self._prepare_connector_prefix_reservation(req) return True def _disagg_transfer_overwrites_whole_cached_prefix(self) -> bool: @@ -3589,8 +3592,8 @@ def _disagg_transfer_overwrites_whole_cached_prefix(self) -> bool: def resize_context(self, req: LlmRequest, num_tokens: int) -> bool: """Resize KV cache to cover context_current_position + num_tokens. - Returns True on success, False if resize failed (first chunk is - suspended on failure). + Return False on allocation failure. Unstarted connector reservations + are released with their tentative allocations. """ assert not req.is_disagg_generation_init_state, ( f"req {req.py_request_id}: use prepare_disagg_gen_init" @@ -3610,7 +3613,15 @@ def resize_context(self, req: LlmRequest, num_tokens: int) -> bool: capacity = max(kv_cache.capacity, target) pre_cap = kv_cache.capacity - success = kv_cache.resize(capacity) + reservation = ( + self.kv_connector_manager.get_prefix_reservation(req) + if self._connector_reservations_enabled() + else None + ) + history = reservation.end if reservation is not None else None + success = ( + kv_cache.resize(capacity, history) if history is not None else kv_cache.resize(capacity) + ) if not success: logger.debug( f"[KVCacheManagerV2] request {req.py_request_id} failed to resize KV cache " @@ -3618,7 +3629,10 @@ def resize_context(self, req: LlmRequest, num_tokens: int) -> bool: f"(context_current_position={req.context_current_position}, " f"num_tokens={num_tokens})" ) - if req.is_first_context_chunk: + if reservation is not None: + self.free_resources(req) + rewind_context_after_cache_drop(req, self.tokens_per_block) + elif req.is_first_context_chunk: kv_cache.suspend() return False self._fill_fresh_kv_pages(req.py_request_id) @@ -3832,11 +3846,9 @@ def prepare_resources(self, scheduled_batch: ScheduledRequests): self._run_kv_connector_hooks(scheduled_batch) def _run_kv_connector_hooks(self, scheduled_batch: ScheduledRequests) -> None: - """Serve the connector prefix, then report the pages it may write into. - - Runs on the batch the forward pass will actually execute; see - ``_apply_connector_matched_prefix`` for why that placement is load-bearing. - """ + """Serve final-batch queries for connectors without source reservations.""" + if self._connector_reservations_enabled(): + return served_any = False for request in scheduled_batch.context_requests: # An allocation is asked about and reported exactly once, and it @@ -3870,27 +3882,99 @@ def _run_kv_connector_hooks(self, scheduled_batch: ScheduledRequests) -> None: # those two lists, so the split has to be rebuilt before it runs. scheduled_batch.reset_context_requests() - def report_batch_to_connector(self, scheduled_batch: ScheduledRequests) -> None: + def report_batch_to_connector( + self, scheduled_batch: ScheduledRequests, *, finalize_prefix_reservations: bool = True + ) -> None: """Report the batch to the KV connector. ``RequestData.num_scheduled_tokens`` describes the upcoming forward pass, so this may only run once every resource manager has. That is why ``ResourceManager.prepare_resources`` drives it rather than ``prepare_resources`` here, which also gives the disagg-generation-init - mini-batch the same hook. + mini-batch the same hook. Such mini-batches must pass + ``finalize_prefix_reservations=False`` to preserve the compute batch's + pending reservations. """ if self.kv_connector_manager is not None and not self.is_draft: + if finalize_prefix_reservations and self._connector_reservations_enabled(): + self._accept_connector_prefix_reservations(scheduled_batch) self.kv_connector_manager.build_scheduler_output(scheduled_batch, self) # ---- KV connector prefix ---- - # - # The connector is asked from `prepare_resources`, downstream of every stage - # that can still drop a request (`_can_queue`, batch waiting, attention-DP - # balancing, the mamba-hybrid filter, the fp8 context-MLA cap). The - # connector ABC has no `cancel_load`, so that placement is required: an - # asked request must always reach `request_finished`. Moving the ask into - # the scheduling pass would buy the scheduler a budget that accounts for the - # served prefix and break the guarantee. + + def _connector_reservations_enabled(self) -> bool: + return ( + self.kv_connector_manager is not None + and not self.is_draft + and self.kv_connector_manager.prefix_reservations_enabled + ) + + def _prepare_connector_prefix_reservation(self, req: LlmRequest) -> None: + """Reserve source KV and expose its prefix to the scheduler's token budget.""" + if ( + not self._connector_reservations_enabled() + or not self._connector_may_serve(req) + or not req.is_first_context_chunk + or req.py_connector_allocation_reported + or not self.kv_connector_manager.should_add_sequence(req) + ): + return + local_end = self.kv_cache_map[req.py_request_id].num_committed_tokens + reservation = self.kv_connector_manager.reserve_prefix(req, local_end) + if reservation is None: + return + # Loads end at full blocks and leave the final prompt token for logits. + end = min(reservation.end, req.prompt_len - 1) + end = end // self.tokens_per_block * self.tokens_per_block + if end <= local_end: + self.kv_connector_manager.release_prefix_reservation(req) + return + reservation = self.kv_connector_manager.trim_prefix_reservation(req, local_end, end) + if reservation is not None: + _settle_context_cursor(req, reservation.end, self.tokens_per_block) + + def release_unused_connector_reservations(self, accepted_request_ids: set[int]) -> None: + """Release unstarted promises and tentative allocations excluded from the batch.""" + if not self._connector_reservations_enabled(): + return + for req in self.kv_connector_manager.pending_prefix_requests(): + if req.request_id in accepted_request_ids: + continue + # Drop unfilled history even when capacity did not grow. + self.free_resources(req) + rewind_context_after_cache_drop(req, self.tokens_per_block) + + def _accept_connector_prefix_reservations(self, scheduled_batch: ScheduledRequests) -> None: + requests = scheduled_batch.context_requests + self.release_unused_connector_reservations({req.request_id for req in requests}) + for req in requests: + if ( + not req.is_first_context_chunk + or req.py_connector_allocation_reported + or not self.kv_connector_manager.should_add_sequence(req) + ): + continue + reservation = self.kv_connector_manager.get_prefix_reservation(req) + by_group = self.get_page_indices_by_layer_group(req) + flat = by_group[0] if len(by_group) == 1 else [] + if reservation is not None: + kv_cache = self.kv_cache_map[req.py_request_id] + allocation_valid = ( + kv_cache.is_active + and kv_cache.history_length >= reservation.end + and req.context_current_position == reservation.end + and kv_cache.capacity >= reservation.end + req.context_chunk_size + ) + if not allocation_valid: + raise RuntimeError( + f"Request {req.request_id} has no allocation for its reserved KV prefix" + ) + self.kv_connector_manager.accept_prefix_load( + req, reservation.start, reservation.end, by_group + ) + req.py_connector_served_position = reservation.end + self.kv_connector_manager.update_state_after_alloc(req, flat, by_group) + req.py_connector_allocation_reported = True def _connector_may_serve(self, req: LlmRequest) -> bool: """Whether the connector is allowed to serve a prefix for ``req``.""" @@ -5075,9 +5159,14 @@ def release_index_slot(self, request_id: int) -> None: self._early_freed_index_requests.add(request_id) def free_resources(self, request: LlmRequest, pin_on_release: bool = False): - # The promise to the connector is one ask per allocation, not one per - # request, so a replay of this request after a destructive pause may be - # asked and reported again. + if self.kv_connector_manager is not None and not self.is_draft: + self.kv_connector_manager.release_unstarted_prefix_loads(request) + if self.kv_connector_manager.has_pending_load(request): + raise RuntimeError( + f"Cannot release KV cache while request {request.request_id} is loading" + ) + self.kv_connector_manager.release_prefix_reservation(request) + # Replay obtains fresh connector metadata for its new destination pages. request.py_connector_allocation_reported = False # Same allocation, same lifetime. The pages the served position vouches # for are gone, so a replay recomputes from its own reuse match. diff --git a/tensorrt_llm/_torch/pyexecutor/llm_request.py b/tensorrt_llm/_torch/pyexecutor/llm_request.py index 299163310c24..a2c7c2548a83 100644 --- a/tensorrt_llm/_torch/pyexecutor/llm_request.py +++ b/tensorrt_llm/_torch/pyexecutor/llm_request.py @@ -1053,16 +1053,12 @@ def __init__( self.py_num_connector_matched_tokens = 0 - # Whether the KV connector has been asked about, and told about, this - # request's current KV allocation. The promise is at most once per - # allocation, not once per request, so this is cleared in - # `free_resources` -- the one place an allocation dies. + # Destination pages are reported once per allocation. Destroying the + # allocation clears this flag so replay can report its new pages. self.py_connector_allocation_reported = False - # End of the prefix a KV connector populated for the current allocation, - # or 0. The cache holds those tokens but never commits them, so a context - # request that re-enters cannot recover the end from the cache's own - # reuse depth. + # End of the connector prefix retained by this allocation, or 0. This + # depth survives async parking; local reuse cannot reconstruct it. self.py_connector_served_position = 0 self.py_result = PyResult( diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 1f0b0d4e78e3..112dab65db48 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -80,7 +80,8 @@ from .disagg_adapter import PyExecutorEffects, PyExecutorRequestRegistry from .dwdp import DwdpManager from .error_classification import ErrorBudget -from .executor_request_queue import (ExecutorRequestQueue, +from .executor_request_queue import (PREFIX_LOAD_COMPLETION_REQUEST_ID, + ExecutorRequestQueue, RequestAdmissionState, RequestQueueItem) from .gpu_keepalive import GpuKeepalive from .guided_decoder import GuidedDecoder @@ -1152,6 +1153,11 @@ def _maybe_init_kv_connector_manager(self): "per-layer load/save hooks have nothing meaningful to " "transfer for those layers.") + self.kv_connector_manager.configure_prefix_reservations( + is_kv_cache_manager_v2 + and (scheduler_config is None + or scheduler_config.enable_prefix_aware_scheduling)) + if is_kv_cache_manager_v2: # Registered regions are device addresses, so every page has # to stay pinned to GPU for as long as the connector holds @@ -1718,6 +1724,8 @@ def shutdown(self): logger.error("Hang detected, shutting down immediately.") return self.worker_thread.join() + if self.kv_connector_manager is not None: + self.kv_connector_manager.shutdown() if self.dist.pp_size > 1: self.executed_batch_queue.put(None) self.broadcast_sample_state_handler.join() @@ -1893,7 +1901,16 @@ def set_gather_responses(self, gather_all_responses): @property def should_stop_processing(self): return self.is_shutdown and len(self.active_requests) == 0 and \ - len(self.waiting_queue) == 0 + len(self.waiting_queue) == 0 and not self._has_pending_connector_transfers() + + def _has_pending_connector_transfers(self) -> bool: + connector = getattr(self, "kv_connector_manager", None) + if connector is None: + return False + transfers = getattr(self, "async_transfer_manager", None) + return (connector.has_pending_loads() + or (transfers is not None + and transfers.has_any_inflight_requests())) @contextmanager def _profiler(self): @@ -3782,6 +3799,7 @@ def _sync_gen_only_benchmark_has_insufficient_kv( return all_ranks_fetched and any_rank_terminal_no_fit def _prepare_and_schedule_batch(self): + self._release_unused_connector_reservations() self._poll_encoder_steps() new_requests = self._fetch_and_activate_new_requests() if self.should_stop_processing: @@ -3930,19 +3948,67 @@ def _prepare_and_schedule_batch(self): f'{scheduled_batch.num_generation_requests} generation requests') return scheduled_batch, iter_stats + def _release_unused_connector_reservations( + self, scheduled_batch: Optional[ScheduledRequests] = None) -> None: + connector = getattr(self, "kv_connector_manager", None) + if connector is None or not connector.prefix_reservations_enabled: + return + accepted = ({ + req.py_request_id + for req in scheduled_batch.context_requests + } if scheduled_batch is not None else set()) + self.kv_cache_manager.release_unused_connector_reservations(accepted) + def _kv_connector_start_batch(self, scheduled_batch): if self.kv_connector_manager: self.kv_connector_manager.take_scheduled_requests_pending_load( scheduled_batch) self.kv_connector_manager.handle_metadata() + self.kv_connector_manager.mark_prefix_loads_dispatched() self.kv_connector_manager.worker.start_load_kv( torch.cuda.current_stream()) + def _defer_connector_load_cancellations(self) -> None: + # A cancelled load must drain without becoming schedulable again. + cancelled = set(self.canceled_req_ids) + for req in self.active_requests: + req_id = (req.parent_request_id + if req.is_child else req.py_request_id) + if req_id in cancelled: + self.kv_connector_manager.defer_load_termination(req) + def _kv_connector_terminate_requests(self): if self.kv_connector_manager: + self._defer_connector_load_cancellations() reqs_to_terminate = self.kv_connector_manager.get_finished() for req in reqs_to_terminate: self._release_transfer(req) + for req in self.kv_connector_manager.take_finished_load_terminations( + ): + self._finish_connector_load_termination(req) + + def _finish_connector_load_termination(self, request: LlmRequest) -> None: + """Finish a cancelled load or reclaim a load whose error was reported.""" + if not request.is_finished: + request.py_kv_transfer_timed_out = False + request.finish_by_reason(FinishReason.CANCELLED) + request.decoding_iter = request.py_decoding_iter + response = request.create_response(False, self.dist.rank) + if response is not None: + response.result.cached_tokens = request.cached_tokens + self._enqueue_responses([(request.py_request_id, response)]) + self.active_requests[:] = [ + req for req in self.active_requests if req is not request + ] + cancel_id = (request.parent_request_id + if request.is_child else request.py_request_id) + if not any((req.parent_request_id if req.is_child else req.py_request_id + ) == cancel_id for req in self.active_requests): + self.canceled_req_ids[:] = [ + req_id for req_id in self.canceled_req_ids + if req_id != cancel_id + ] + self._terminate_request(request) def _kv_connector_wait_for_save(self): if self.kv_connector_manager is not None: @@ -4289,6 +4355,7 @@ def _executor_loop(self): can_forward, should_retry = self._check_benchmark_disagg_gate( scheduled_batch, can_forward) if should_retry: + self._release_unused_connector_reservations() if self._is_kv_manager_v2: self._revert_gen_alloc(scheduled_batch) self._terminate_recompute_paused_requests( @@ -4321,6 +4388,8 @@ def _executor_loop(self): gpu_forward_events_from_perf_pool = False can_queue, _ = self._can_queue(scheduled_batch) + self._release_unused_connector_reservations( + scheduled_batch if can_queue else None) if can_queue: self._prepare_disagg_gen_transmission_complete( @@ -4571,8 +4640,9 @@ def _handle_control_request(self): pending = self.control_requests[0] - if pending.control_requires_drain and (len(self.active_requests) != 0 - or len(self.waiting_queue) != 0): + if pending.control_requires_drain and ( + len(self.active_requests) != 0 or len(self.waiting_queue) != 0 + or self._has_pending_connector_transfers()): # drain=True: keep the sentinel parked until the engine drains. return @@ -5129,6 +5199,7 @@ def _executor_loop_overlap(self): can_forward, should_retry = self._check_benchmark_disagg_gate( scheduled_batch, can_forward) if should_retry: + self._release_unused_connector_reservations() if self._is_kv_manager_v2: self._revert_gen_alloc(scheduled_batch) self._terminate_recompute_paused_requests( @@ -5161,6 +5232,8 @@ def _executor_loop_overlap(self): can_queue, can_queue_this_rank = self._can_queue( scheduled_batch) + self._release_unused_connector_reservations( + scheduled_batch if can_queue else None) if can_queue: self._prepare_disagg_gen_transmission_complete( @@ -5719,8 +5792,12 @@ def _validate_request(self, request: LlmRequest): def _fetch_and_enqueue_requests(self, waiting_queue: WaitingQueue, total_num_live_requests: int) -> None: """Fetch requests from request_queue and enqueue to waiting_queue.""" - # Block new requests while control requests are pending - if len(self.control_requests) != 0: + connector = self.kv_connector_manager + control_pending = len(self.control_requests) != 0 + # A draining control request must still receive transfer completions. + if control_pending and not (connector is not None + and connector.prefix_reservations_enabled + and connector.has_pending_loads()): return # Calculate timeout. Never wait once the shutdown sentinel has been @@ -5729,7 +5806,8 @@ def _fetch_and_enqueue_requests(self, waiting_queue: WaitingQueue, # `should_stop_processing` check that ends it, deadlocking shutdown() # on `shutdown_event`. idle = (total_num_live_requests == 0 and len(waiting_queue) == 0 - and not self.is_shutdown) + and not self.is_shutdown + and not self._has_pending_connector_transfers()) if idle: # In Ray path (TLLM_DISABLE_MPI=1), use a periodic heartbeat timeout so rank 0 # reaches the broadcast path regularly to prevent trtllm-serve timeout when idle. @@ -5740,7 +5818,7 @@ def _fetch_and_enqueue_requests(self, waiting_queue: WaitingQueue, # Fetch requests from rank 0 new_requests = [] - if self.dist.rank == 0: + if self.dist.rank == 0 and not control_pending: # Process accumulated requests that were queued during control request handling. if len(self.request_accumulated) != 0: new_requests.extend(self.request_accumulated) @@ -5751,6 +5829,16 @@ def _fetch_and_enqueue_requests(self, waiting_queue: WaitingQueue, new_requests.extend( self.executor_request_queue.get_from_request_queue(timeout)) + if connector is not None and self.dist.rank == 0: + completed = connector.take_finished_prefix_loads() + if completed: + # Apply completions even when a shutdown or control item stops + # consumption of the rest of this request envelope. + new_requests.insert( + 0, + RequestQueueItem(PREFIX_LOAD_COMPLETION_REQUEST_ID, + finished_prefix_load_ids=completed)) + # Broadcast requests and handle Python objects. RequestBroadcaster probes # the request count first and can skip the heavy payload broadcast on # empty iterations. @@ -5995,8 +6083,12 @@ def _handle_special_queue_items( new_requests: List[RequestQueueItem]) -> List[RequestQueueItem]: """Handle special signals.""" accepted_new_requests = [] + finished_prefix_load_ids = [] for idx, req_item in enumerate(new_requests): - if req_item.is_shutdown_request: + if req_item.is_prefix_load_completion_request: + finished_prefix_load_ids.extend( + req_item.finished_prefix_load_ids) + elif req_item.is_shutdown_request: self.is_shutdown = True break elif req_item.is_canceled_request: @@ -6017,6 +6109,10 @@ def _handle_special_queue_items( else: accepted_new_requests.append(req_item) + if finished_prefix_load_ids: + self._defer_connector_load_cancellations() + self.kv_connector_manager.finish_prefix_loads( + finished_prefix_load_ids) return accepted_new_requests def _update_new_active_requests_queue_latency( @@ -7183,7 +7279,10 @@ def _prepare_disagg_gen_resources(self, requests: List[LlmRequest]) -> None: # so the connector's per-request state advances exactly as before. kv_cache_manager = self.resource_manager.resource_managers.get( ResourceManagerType.KV_CACHE_MANAGER) - if hasattr(kv_cache_manager, "report_batch_to_connector"): + if isinstance(kv_cache_manager, KVCacheManagerV2): + kv_cache_manager.report_batch_to_connector( + disagg_gen_init_to_prepare, finalize_prefix_reservations=False) + elif hasattr(kv_cache_manager, "report_batch_to_connector"): kv_cache_manager.report_batch_to_connector( disagg_gen_init_to_prepare) @@ -7971,6 +8070,15 @@ def _handle_errors(self, raise self._fatal_error def _terminate_request(self, request: LlmRequest) -> None: + connector = getattr(self, "kv_connector_manager", None) + if connector is not None: + connector.release_unstarted_prefix_loads(request) + if connector.defer_load_termination(request): + return + transfers = getattr(self, "async_transfer_manager", None) + if (self._is_kv_manager_v2 and transfers is not None and + request.py_request_id in transfers.requests_in_transfer()): + return # Dummy requests don't participate in disagg KV cache transfers, # so they must bypass the PP termination handler to avoid stale # sequences in the KV cache manager (the handler delays removal, @@ -8037,6 +8145,11 @@ def _try_cancel_request(self, request) -> bool: Returns: bool: True if the request can be canceled (either successfully cancelled or doesn't need cancellation). """ + connector = getattr(self, "kv_connector_manager", None) + if connector is not None: + connector.release_unstarted_prefix_loads(request) + if connector.defer_load_termination(request): + return False if self.kv_cache_transceiver is None: return True diff --git a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py index b49877c6a68c..2929fe386200 100644 --- a/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py +++ b/tensorrt_llm/_torch/pyexecutor/scheduler/scheduler_v2.py @@ -533,12 +533,14 @@ def _schedule_loop(self, active_requests, inflight_request_ids): if r.is_generation_in_progress_state and not r.is_generation_to_complete_state and r.request_id not in inflight_request_ids + and not self._has_pending_connector_load(r) ) if ( num_gen_candidates > 0 and not evicted and not recompute_paused and not inflight_request_ids + and not any(self._has_pending_connector_load(r) for r in active_requests) ): # A connector rejects every tier below GPU at bring-up # (`PyExecutor._reject_non_gpu_cache_tiers`), so offering those @@ -739,17 +741,9 @@ def _try_schedule_context( Returns ``(action, tokens, chunking_flag)``. *tokens* and *chunking_flag* are meaningful only when *action* is ``SCHEDULED``. - Deliberately connector-blind: the connector is asked later, in - the cache manager's ``prepare_resources``, so the budget here assumes it - serves nothing. That is safe because honouring an offer only ever - removes tokens from the forward pass, and it costs only that a served - prefix frees no budget for another request in the same iteration. - - ``should_add_sequence`` must not gate scheduling either: it stays false - from the moment an asynchronous load completes until - ``request_finished``, so a request gated on it would be skipped forever - and never run the prefill the load was for. A loading request is kept - out of the batch by its ``DISAGG_GENERATION_TRANS_IN_PROGRESS`` state. + Source reservations advance the prefix before budget checks. Destination + pages are allocated here; loads are accepted after the final batch trims. + A loading request remains outside the schedulable state range. """ first_chunk = req.is_first_context_chunk if self.chunking_enabled: @@ -1322,7 +1316,9 @@ def _try_schedule_generation( # GPU pages so other requests can resume(). # Skip if already suspended — suspending again is a no-op # that frees no pages. - if self.kv_cache_manager.is_request_active(req.py_request_id): + if self.kv_cache_manager.is_request_active( + req.py_request_id + ) and not self._has_pending_connector_load(req): logger.debug( f"[V2Scheduler] Self-evicting request {req.py_request_id} " f"(state={req.state.name}) to free GPU pages" @@ -1384,6 +1380,10 @@ def _free_kv_caches(self, req: LlmRequest) -> None: if self.draft_kv_cache_manager is not None: self.draft_kv_cache_manager.free_resources(req) + def _has_pending_connector_load(self, req: LlmRequest) -> bool: + connector = self.kv_cache_manager.kv_connector_manager + return connector is not None and connector.has_pending_load(req) + def _is_evictable(self, req: LlmRequest, inflight_request_ids: set[int]) -> bool: """A started request whose KV cache is still active on GPU. @@ -1394,6 +1394,8 @@ def _is_evictable(self, req: LlmRequest, inflight_request_ids: set[int]) -> bool return False if not self._is_started_request(req): return False + if self._has_pending_connector_load(req): + return False return self.kv_cache_manager.is_request_active(req.py_request_id) def _is_recompute_pause_candidate( @@ -1401,6 +1403,8 @@ def _is_recompute_pause_candidate( ) -> bool: if req.request_id in inflight_request_ids: return False + if self._has_pending_connector_load(req): + return False # is_generation_in_progress_state also includes GENERATION_TO_COMPLETE, # which is outside the schedulable range and may still be finalizing. if req.state_value == self._gen_to_complete_state_value: diff --git a/tests/integration/defs/llmapi/test_llm_api_connector.py b/tests/integration/defs/llmapi/test_llm_api_connector.py index aac0ade227b4..141b6ed6af3d 100644 --- a/tests/integration/defs/llmapi/test_llm_api_connector.py +++ b/tests/integration/defs/llmapi/test_llm_api_connector.py @@ -21,13 +21,16 @@ import sys import tempfile import time +from threading import Event +from types import SimpleNamespace from unittest.mock import MagicMock, patch import pytest from tensorrt_llm import LLM, DisaggregatedParams, SamplingParams from tensorrt_llm._torch.pyexecutor.connectors.kv_cache_connector import ( - V2_RETENTION_IGNORED_LOG_KEY, KvCacheConnectorWorker) + V2_RETENTION_IGNORED_LOG_KEY, KvCacheConnectorManager, + KvCacheConnectorWorker, PrefixLoad, SchedulerOutput) from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import \ KVCacheManagerV2 from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager @@ -1627,8 +1630,8 @@ def test_connector_e2e_persistent_cache(enforce_single_worker, # to the root logger, so read it from the return value instead. matched_tokens = [] leader_cls = llm_kv_cache_connector.PersistentKvCacheConnectorLeader - original_get_num_new_matched_tokens = ( - leader_cls.get_num_new_matched_tokens) + original_get_num_new_matched_tokens = leader_cls.get_num_new_matched_tokens + original_reserve_prefix = leader_cls.reserve_prefix def recording_get_num_new_matched_tokens(self, request, num_computed_tokens): @@ -1637,8 +1640,17 @@ def recording_get_num_new_matched_tokens(self, request, matched_tokens.append(result[0]) return result + def recording_reserve_prefix(self, request, num_computed_tokens, + reservation_id): + result = original_reserve_prefix(self, request, num_computed_tokens, + reservation_id) + matched_tokens.append(result[0]) + return result + monkeypatch.setattr(leader_cls, "get_num_new_matched_tokens", recording_get_num_new_matched_tokens) + monkeypatch.setattr(leader_cls, "reserve_prefix", + recording_reserve_prefix) kv_connector_config = KvCacheConnectorConfig( connector_module="llm_kv_cache_connector", @@ -1738,6 +1750,251 @@ def recording_get_num_new_matched_tokens(self, request, shutil.rmtree(cache_dir, ignore_errors=True) +def test_connector_reservations_hold_source_without_transmission( + monkeypatch, tmp_path): + """Range release and cache-key removal preserve another promised source.""" + import torch + + examples_dir = os.path.abspath( + os.path.join(os.path.dirname(__file__), "..", "..", "..", "..", + "examples", "llm-api")) + monkeypatch.syspath_prepend(examples_dir) + monkeypatch.setenv("CONNECTOR_CACHE_FOLDER", str(tmp_path)) + import llm_kv_cache_connector as persistent + + args = SimpleNamespace(kv_cache_config=SimpleNamespace(tokens_per_block=4)) + leader = persistent.PersistentKvCacheConnectorLeader(args) + request = SimpleNamespace(request_id=7, + cache_salt="test", + get_tokens=lambda beam: [1, 2, 3, 4, 5]) + key = leader._hash_tokens([1, 2, 3, 4], request.cache_salt) + source_path = leader._file_path(key) + source = torch.arange(4) + torch.save(source, source_path) + + with patch.object(torch, + "load", + side_effect=AssertionError("Premature read")): + assert leader.reserve_prefix(request, 0, 11) == (4, False) + assert leader.reserve_prefix(request, 0, 12) == (4, False) + source_path.unlink() + torch.save(source + 100, source_path) + leader.release_prefix_reservation(request, 11, 0, 2) + assert leader._reserved_files[11][0].exists() + leader.release_prefix_reservation(request, 11, 2, 4) + assert 11 not in leader._reserved_files + assert leader._reserved_files[12][0].exists() + metadata = leader.build_connector_meta( + SchedulerOutput(prefix_loads=[ + PrefixLoad(reservation_id=12, + request_id=request.request_id, + start=0, + end=4, + is_async=False, + block_ids_by_layer_group=[[1]], + tokens=[1, 2, 3, 4, 5], + cache_salt=request.cache_salt) + ])) + + worker = persistent.PersistentKvCacheConnectorWorker(args) + destination = torch.zeros((2, 4), dtype=source.dtype) + worker.register_kv_caches(destination) + worker._metadata = metadata + assert worker.get_finished_prefix_loads() == [] + worker.start_load_kv(None) + assert torch.equal(destination[1], source) + assert not destination[0].any() + assert worker.get_finished_prefix_loads() == [12] + assert worker.get_finished_prefix_loads() == [] + leader.release_prefix_reservation(request, 12, 0, 4) + assert not leader._reserved_files + assert not leader._reserved_ranges + + +@pytest.fixture +def controlled_prefix_connector(enforce_single_worker, monkeypatch, tmp_path): + """Keep real disk-to-device writes pending until the test releases them.""" + examples_dir = os.path.abspath( + os.path.join(os.path.dirname(__file__), "..", "..", "..", "..", + "examples", "llm-api")) + monkeypatch.syspath_prepend(examples_dir) + monkeypatch.setenv("CONNECTOR_CACHE_FOLDER", str(tmp_path)) + import llm_kv_cache_connector as persistent + import torch + + control = SimpleNamespace(started=Event(), + allow_copy=Event(), + copied=Event(), + cancellation_seen=Event(), + freed=Event(), + source_released=Event(), + request_id=None, + accepted=None, + copies=0, + compute_ids=[]) + + class ControlledLeader(persistent.PersistentKvCacheConnectorLeader): + + def reserve_prefix(self, request, num_computed_tokens, reservation_id): + count, _ = super().reserve_prefix(request, num_computed_tokens, + reservation_id) + return count, bool(count) + + def build_connector_meta(self, scheduler_output): + control.compute_ids.extend( + req.request_id for req in scheduler_output.new_requests + + scheduler_output.cached_requests) + for load in scheduler_output.prefix_loads: + assert control.accepted is None, "Load dispatched twice" + control.accepted = load + control.request_id = load.request_id + return super().build_connector_meta(scheduler_output) + + def release_prefix_reservation(self, request, reservation_id, start, + end): + accepted = control.accepted + if (accepted is not None + and reservation_id == accepted.reservation_id + and start < accepted.end and end > accepted.start): + assert control.copied.is_set(), "Source released before copy" + control.source_released.set() + super().release_prefix_reservation(request, reservation_id, start, + end) + + class ControlledWorker(persistent.PersistentKvCacheConnectorWorker): + + def __init__(self, llm_args): + super().__init__(llm_args) + self.pending = None + + def start_load_kv(self, stream): + if not self._metadata.prefix_load_ids: + return + assert self.pending is None, "Load dispatched twice" + self.pending = self._metadata + assert control.accepted is not None + control.started.set() + + def get_finished_prefix_loads(self): + if self.pending is None or not control.allow_copy.is_set(): + return [] + for path, slot in self.pending.load: + source = torch.load(path, map_location="cpu") + self.kv_cache_tensor[slot].copy_(source, non_blocking=False) + assert torch.equal(self.kv_cache_tensor[slot].cpu(), source) + control.copies += 1 + finished = self.pending.prefix_load_ids + self.pending = None + control.copied.set() + return finished + + original_defer = KvCacheConnectorManager.defer_load_termination + original_free = KVCacheManagerV2.free_resources + + def record_defer(manager, request): + deferred = original_defer(manager, request) + if deferred and request.request_id == control.request_id: + control.cancellation_seen.set() + return deferred + + def record_free(manager, request): + target = request.request_id == control.request_id + if target: + assert control.copied.is_set(), "Destination freed before copy" + result = original_free(manager, request) + if target: + control.freed.set() + return result + + monkeypatch.setattr(persistent, + "ControlledLeader", + ControlledLeader, + raising=False) + monkeypatch.setattr(persistent, + "ControlledWorker", + ControlledWorker, + raising=False) + monkeypatch.setattr(KvCacheConnectorManager, "defer_load_termination", + record_defer) + monkeypatch.setattr(KVCacheManagerV2, "free_resources", record_free) + yield persistent, control + control.allow_copy.set() + + +@pytest.mark.threadleak(enabled=False) +@pytest.mark.parametrize("use_overlap_scheduler", [True, False]) +@pytest.mark.parametrize("other_traffic", [False, True], + ids=["alone", "traffic"]) +def test_connector_cancel_drains_real_prefix_load(controlled_prefix_connector, + use_overlap_scheduler, + other_traffic): + """Cancellation retains source and destination until real writes finish.""" + persistent, control = controlled_prefix_connector + kwargs = dict( + model=f"{llm_models_root()}/Qwen3/Qwen3-0.6B", + backend="pytorch", + cuda_graph_config=None, + disable_overlap_scheduler=not use_overlap_scheduler, + max_seq_len=256, + max_num_tokens=128, + max_batch_size=2, + enable_chunked_prefill=False, + scheduler_config=SchedulerConfig(enable_prefix_aware_scheduling=True), + kv_cache_config=KvCacheConfig(max_tokens=192, + tokens_per_block=32, + use_kv_cache_manager_v2=True), + ) + params = SamplingParams(max_tokens=2, ignore_eos=True) + prompt = [100] * 96 + cold = LLM(**kwargs, + kv_connector_config=KvCacheConnectorConfig( + connector_module=persistent.__name__, + connector_scheduler_class="PersistentKvCacheConnectorLeader", + connector_worker_class="PersistentKvCacheConnectorWorker")) + try: + cold.generate(prompt, params) + finally: + cold.shutdown() + + warm = LLM(**kwargs, + kv_connector_config=KvCacheConnectorConfig( + connector_module=persistent.__name__, + connector_scheduler_class="ControlledLeader", + connector_worker_class="ControlledWorker")) + try: + cancelled = warm.generate_async(prompt, params) + assert control.started.wait( + 60), "Confirmed async load was not dispatched" + assert control.accepted.is_async + assert control.accepted.end > control.accepted.start + assert control.copies == 0 + cancelled.abort() + assert control.cancellation_seen.wait( + 60), "Cancellation was not consumed" + assert not control.freed.is_set() + assert not control.source_released.is_set() + + if other_traffic: + other = warm.generate_async([200] * 32, params) + assert len(other.result(timeout=60).outputs[0].token_ids) == 2 + assert not control.freed.is_set() + + control.allow_copy.set() + assert control.freed.wait( + 60), "Cancelled load never released its allocation" + assert control.copied.is_set() + assert control.source_released.is_set() + assert control.copies > 0 + assert control.request_id not in control.compute_ids + + # This request needs more capacity than remains while the load owns KV. + after = warm.generate_async([300] * 128, params) + assert len(after.result(timeout=60).outputs[0].token_ids) == 2 + finally: + control.allow_copy.set() + warm.shutdown() + + # The VSWA end-to-end sizes. The window is deliberately larger than the whole # run (prompt + generation), so nothing goes out of window and the save/load # round trip is the only thing under test. `test_connector_vswa_reports_page_ diff --git a/tests/integration/test_lists/test-db/l0_a10.yml b/tests/integration/test_lists/test-db/l0_a10.yml index 8415c70c1dab..b35205380b5f 100644 --- a/tests/integration/test_lists/test-db/l0_a10.yml +++ b/tests/integration/test_lists/test-db/l0_a10.yml @@ -224,6 +224,11 @@ l0_a10: - llmapi/test_llm_api_connector.py::test_connector_max_utilization_is_rejected_on_v1_only[kv_cache_manager_v1] - llmapi/test_llm_api_connector.py::test_connector_max_utilization_is_rejected_on_v1_only[kv_cache_manager_v2] - llmapi/test_llm_api_connector.py::test_connector_e2e_persistent_cache[kv_cache_manager_v2] + - llmapi/test_llm_api_connector.py::test_connector_reservations_hold_source_without_transmission + - llmapi/test_llm_api_connector.py::test_connector_cancel_drains_real_prefix_load[alone-True] + - llmapi/test_llm_api_connector.py::test_connector_cancel_drains_real_prefix_load[alone-False] + - llmapi/test_llm_api_connector.py::test_connector_cancel_drains_real_prefix_load[traffic-True] + - llmapi/test_llm_api_connector.py::test_connector_cancel_drains_real_prefix_load[traffic-False] - llmapi/test_llm_api_connector.py::test_connector_prefix_is_asked_once_and_shrinks_the_forward_pass[kv_cache_manager_v2] - llmapi/test_llm_api_connector.py::test_connector_prefix_under_chunked_prefill[kv_cache_manager_v2-offer_inside_chunk] - llmapi/test_llm_api_connector.py::test_connector_prefix_under_chunked_prefill[kv_cache_manager_v2-offer_past_chunk] diff --git a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py index 4124d7b657c7..7a941a960796 100644 --- a/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py +++ b/tests/unittest/_torch/executor/kv_cache/test_kv_cache_v2_scheduler.py @@ -174,6 +174,7 @@ def make_kv_cache_manager( is_vswa=False, ): mgr = Mock() + mgr.kv_connector_manager = None mgr.tokens_per_block = tokens_per_block # A real policy, not an auto-created Mock attribute: the prefix-aware skip # is gated on ALL_REUSABLE (see _skip_pays_off_under_reuse_policy). @@ -3456,3 +3457,110 @@ def test_recompute_paused_request_does_not_defer_duplicate(self): # on behalf of a request that this iteration paused. assert 1 not in ids(out.paused_requests) assert ids(out.context_requests) in ([1], []) + + +@pytest.mark.parametrize("chunked", [False, True]) +def test_connector_credit_admits_another_request(chunked: bool) -> None: + from types import SimpleNamespace + + from tensorrt_llm._torch.pyexecutor.kv_cache.kv_cache_manager_v2 import KVCacheManagerV2 + from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, SamplingConfig + from tensorrt_llm._torch.pyexecutor.scheduler.scheduler_v2 import BudgetTracker + + def run_with_offer(offered_tokens: int) -> tuple[object, list[int], Mock]: + mgr = make_kv_cache_manager(tokens_per_block=32, enable_block_reuse=True) + mgr.is_draft = False + connector = Mock(prefix_reservations_enabled=True) + connector.should_add_sequence.return_value = True + connector.reserve_prefix.side_effect = lambda req, local_end: ( + SimpleNamespace(start=local_end, end=local_end + offered_tokens) + if offered_tokens + else None + ) + connector.trim_prefix_reservation.side_effect = lambda req, start, end: SimpleNamespace( + start=start, end=end + ) + mgr.kv_connector_manager = connector + for name in ( + "_connector_reservations_enabled", + "_connector_may_serve", + "_prepare_connector_prefix_reservation", + ): + setattr(mgr, name, getattr(KVCacheManagerV2, name).__get__(mgr)) + mgr.prepare_context.side_effect = KVCacheManagerV2.prepare_context.__get__(mgr) + requests = [ + LlmRequest( + request_id=request_id, + max_new_tokens=1, + input_tokens=list(range(request_id * 100, request_id * 100 + 96)), + sampling_config=SamplingConfig(1), + is_streaming=False, + ) + for request_id in (1, 2) + ] + for req in requests: + mgr.kv_cache_map[req.request_id].num_committed_tokens = 0 + scheduler = make_scheduler( + mgr, + max_num_tokens=64 if chunked else 128, + ctx_chunk_config=(None, 32) if chunked else None, + ) + original_commit = BudgetTracker.commit + with patch.object( + BudgetTracker, "commit", autospec=True, side_effect=original_commit + ) as commit: + output = scheduler.schedule_request(requests, set()) + charges = [entry.args[2] for entry in commit.call_args_list] + connector.accept_prefix_load.assert_not_called() + return output, charges, mgr + + baseline, baseline_charges, _ = run_with_offer(0) + credited, credited_charges, manager = run_with_offer(64) + + assert ids(baseline.context_requests) == [1] + assert baseline_charges == [64 if chunked else 96] + assert ids(credited.context_requests) == [1, 2] + assert credited_charges == [32, 32] + assert [call.args[1] for call in manager.resize_context.call_args_list] == [32, 32] + + +def test_pending_connector_load_is_not_a_recompute_or_eviction_victim() -> None: + manager = make_kv_cache_manager(can_evict=True) + connector = Mock() + connector.has_pending_load.return_value = True + manager.kv_connector_manager = connector + request = make_gen_request(1) + scheduler = make_scheduler(manager) + assert not scheduler._is_evictable(request, set()) + assert not scheduler._is_recompute_pause_candidate(request, set()) + + +def test_pending_connector_load_does_not_self_evict_or_report_deadlock() -> None: + manager = make_kv_cache_manager( + can_evict=True, try_allocate_generation_fn=lambda request: False + ) + connector = Mock() + connector.has_pending_load.return_value = True + manager.kv_connector_manager = connector + request = make_gen_request(1) + scheduler = make_scheduler(manager) + output = scheduler.schedule_request([request], set()) + assert output.generation_requests == [] + assert output.paused_requests == [] + assert output.recompute_paused_requests == [] + manager.suspend_request.assert_not_called() + manager.free_resources.assert_not_called() + + +def test_parked_connector_load_keeps_kv_pressure_retryable() -> None: + manager = make_kv_cache_manager(try_allocate_generation_fn=lambda request: False) + connector = Mock() + connector.has_pending_load.side_effect = lambda request: request.request_id == 1 + manager.kv_connector_manager = connector + loading = make_filtered_request(1, state_value=DISAGG_GEN_TRANS_IN_PROGRESS) + generation = make_gen_request(2) + manager.kv_cache_map[2].is_active = False + scheduler = make_scheduler(manager) + output = scheduler.schedule_request([loading, generation], set()) + assert output.generation_requests == [] + assert output.recompute_paused_requests == [] diff --git a/tests/unittest/_torch/executor/test_kv_connector_executor_lifetime.py b/tests/unittest/_torch/executor/test_kv_connector_executor_lifetime.py new file mode 100644 index 000000000000..eb8348da3c33 --- /dev/null +++ b/tests/unittest/_torch/executor/test_kv_connector_executor_lifetime.py @@ -0,0 +1,393 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""Executor ownership of connector loads through cancellation and completion.""" + +from contextlib import nullcontext +from copy import deepcopy +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest + +from tensorrt_llm._torch.pyexecutor.connectors.kv_cache_connector import KvCacheConnectorManager +from tensorrt_llm._torch.pyexecutor.executor_request_queue import ( + CONTROL_REQUEST_ID, + PREFIX_LOAD_COMPLETION_REQUEST_ID, + SHUTDOWN_REQUEST_ID, + RequestQueueItem, +) +from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor +from tensorrt_llm._torch.pyexecutor.request_utils import RequestBroadcaster +from tensorrt_llm.bindings.executor import FinishReason + +pytestmark = pytest.mark.cpu_only + + +def _executor() -> PyExecutor: + executor = object.__new__(PyExecutor) + executor.kv_connector_manager = Mock(spec=KvCacheConnectorManager) + executor.kv_connector_manager.prefix_reservations_enabled = True + executor.kv_connector_manager.defer_load_termination.return_value = False + executor.kv_connector_manager.has_pending_loads.return_value = False + executor.kv_connector_manager.get_finished.return_value = [] + executor.kv_connector_manager.take_finished_load_terminations.return_value = [] + executor.kv_cache_manager = Mock() + executor._is_kv_manager_v2 = True + executor.kv_cache_transceiver = None + executor.active_requests = [] + executor.canceled_req_ids = [] + executor.dist = SimpleNamespace(rank=0) + executor._disagg_pp_termination_handler = None + executor._do_terminate_request = Mock() + executor._enqueue_responses = Mock() + executor._release_transfer = Mock() + return executor + + +def _request(request_id: int = 1) -> SimpleNamespace: + response = SimpleNamespace(result=SimpleNamespace(cached_tokens=0)) + request = SimpleNamespace( + py_request_id=request_id, + is_child=False, + is_dummy_request=False, + is_finished=False, + py_decoding_iter=0, + cached_tokens=32, + create_response=Mock(return_value=response), + ) + + def finish(reason: FinishReason) -> None: + assert reason == FinishReason.CANCELLED + request.is_finished = True + + request.finish_by_reason = Mock(side_effect=finish) + return request + + +def test_cancel_waits_for_connector_without_a_transceiver() -> None: + executor = _executor() + request = _request() + executor.kv_connector_manager.defer_load_termination.return_value = True + + assert not executor._try_cancel_request(request) + executor.kv_connector_manager.defer_load_termination.assert_called_once_with(request) + assert not request.is_finished + executor._do_terminate_request.assert_not_called() + + +def test_teardown_cannot_free_an_outstanding_load() -> None: + executor = _executor() + request = _request() + executor.kv_connector_manager.defer_load_termination.return_value = True + + executor._terminate_request(request) + + executor.kv_connector_manager.release_unstarted_prefix_loads.assert_called_once_with(request) + executor.kv_connector_manager.defer_load_termination.assert_called_once_with(request) + executor._do_terminate_request.assert_not_called() + + +def test_unstarted_load_can_be_released_before_teardown() -> None: + executor = _executor() + request = _request() + calls = [] + executor.kv_connector_manager.release_unstarted_prefix_loads.side_effect = ( + lambda req: calls.append("release") + ) + executor._do_terminate_request.side_effect = lambda req: calls.append("free") + + executor._terminate_request(request) + + assert calls == ["release", "free"] + + +def test_lone_loading_request_consumes_cancellation_before_completion() -> None: + executor = _executor() + request = _request() + executor.active_requests = [request] + executor.canceled_req_ids = [request.py_request_id] + connector = executor.kv_connector_manager + calls = [] + connector.defer_load_termination.side_effect = lambda req: calls.append("defer") or False + + def complete() -> list: + assert calls == ["defer"] + calls.append("complete") + return [] + + connector.get_finished.side_effect = complete + connector.take_finished_load_terminations.return_value = [request] + + executor._kv_connector_terminate_requests() + + request.finish_by_reason.assert_called_once_with(FinishReason.CANCELLED) + assert executor.active_requests == [] + assert executor.canceled_req_ids == [] + response = request.create_response.return_value + assert response.result.cached_tokens == 32 + executor._enqueue_responses.assert_called_once_with([(request.py_request_id, response)]) + executor._do_terminate_request.assert_called_once_with(request) + executor._release_transfer.assert_not_called() + + +def test_load_error_already_reported_does_not_emit_another_response() -> None: + executor = _executor() + request = _request() + request.is_finished = True + executor.kv_connector_manager.take_finished_load_terminations.return_value = [request] + + executor._kv_connector_terminate_requests() + + request.create_response.assert_not_called() + executor._enqueue_responses.assert_not_called() + executor._do_terminate_request.assert_called_once_with(request) + + +def test_allocation_survives_until_the_completion_poll() -> None: + executor = _executor() + request = _request() + executor.active_requests = [request] + executor.canceled_req_ids = [request.py_request_id] + executor.kv_connector_manager.defer_load_termination.return_value = True + + executor._kv_connector_terminate_requests() + + assert executor.active_requests == [request] + assert executor.canceled_req_ids == [request.py_request_id] + request.finish_by_reason.assert_not_called() + executor._do_terminate_request.assert_not_called() + + +def test_load_dispatch_marks_ownership_before_worker_launch(monkeypatch) -> None: + executor = _executor() + connector = executor.kv_connector_manager + calls = [] + connector.take_scheduled_requests_pending_load.side_effect = lambda batch: calls.append("park") + connector.handle_metadata.side_effect = lambda: calls.append("metadata") + connector.mark_prefix_loads_dispatched.side_effect = lambda: calls.append("own") + connector.worker = Mock() + connector.worker.start_load_kv.side_effect = lambda stream: calls.append("start") + monkeypatch.setattr("torch.cuda.current_stream", lambda: None) + + executor._kv_connector_start_batch(SimpleNamespace()) + + assert calls == ["park", "metadata", "own", "start"] + + +@pytest.mark.parametrize("selected", [True, False]) +def test_final_batch_releases_every_unselected_reservation(selected: bool) -> None: + executor = _executor() + batch = SimpleNamespace(context_requests=[_request(7)]) if selected else None + + executor._release_unused_connector_reservations(batch) + + executor.kv_cache_manager.release_unused_connector_reservations.assert_called_once_with( + {7} if selected else set() + ) + + +def test_shutdown_waits_for_an_outstanding_load_after_request_error() -> None: + executor = _executor() + executor.is_shutdown = True + executor.waiting_queue = [] + executor.kv_connector_manager.has_pending_loads.return_value = True + + assert not executor.should_stop_processing + + executor.kv_connector_manager.has_pending_loads.return_value = False + assert executor.should_stop_processing + + +def test_completing_one_child_preserves_cancellation_for_its_sibling() -> None: + executor = _executor() + loading = _request(1) + waiting = _request(2) + for request in (loading, waiting): + request.is_child = True + request.parent_request_id = 9 + executor.active_requests = [loading, waiting] + executor.canceled_req_ids = [9] + + executor._finish_connector_load_termination(loading) + + assert executor.active_requests == [waiting] + assert executor.canceled_req_ids == [9] + executor._do_terminate_request.assert_called_once_with(loading) + + +def test_completed_load_does_not_release_a_pending_save() -> None: + executor = _executor() + request = _request() + request.is_finished = True + executor.async_transfer_manager = Mock() + executor.async_transfer_manager.requests_in_transfer.return_value = { + request.py_request_id: request + } + + executor._finish_connector_load_termination(request) + + executor._do_terminate_request.assert_not_called() + executor._enqueue_responses.assert_not_called() + + executor.async_transfer_manager.requests_in_transfer.return_value = {} + executor._terminate_request(request) + executor._do_terminate_request.assert_called_once_with(request) + + +def test_polling_continues_for_save_after_load_completion() -> None: + executor = _executor() + executor.is_shutdown = True + executor.waiting_queue = [] + executor.async_transfer_manager = Mock() + executor.async_transfer_manager.has_any_inflight_requests.return_value = True + + assert not executor.should_stop_processing + + executor.async_transfer_manager.has_any_inflight_requests.return_value = False + assert executor.should_stop_processing + + +def test_legacy_pinned_save_keeps_its_existing_early_teardown() -> None: + executor = _executor() + executor._is_kv_manager_v2 = False + request = _request() + request.is_finished = True + executor.async_transfer_manager = Mock() + executor.async_transfer_manager.requests_in_transfer.return_value = { + request.py_request_id: request + } + + executor._terminate_request(request) + + executor._do_terminate_request.assert_called_once_with(request) + + +def test_control_drain_waits_for_detached_connector_load() -> None: + executor = _executor() + executor.waiting_queue = [] + pending = SimpleNamespace(control_requires_drain=True) + executor.control_requests = [pending] + executor.kv_connector_manager.has_pending_loads.return_value = True + + executor._handle_control_request() + + assert executor.control_requests == [pending] + + +@pytest.mark.parametrize("control_pending", [False, True]) +def test_completion_is_broadcast_without_new_requests(control_pending: bool) -> None: + executor = _executor() + connector = executor.kv_connector_manager + connector.has_pending_loads.return_value = True + connector.take_finished_prefix_loads.return_value = [17] + executor.control_requests = [RequestQueueItem(CONTROL_REQUEST_ID)] if control_pending else [] + executor.is_shutdown = False + executor._disable_mpi = False + executor.request_accumulated = [] + executor.hang_detector = SimpleNamespace(pause=nullcontext) + executor.dist.world_size = 1 + executor.request_broadcaster = RequestBroadcaster(executor.dist, executor.hang_detector) + executor.executor_request_queue = Mock() + executor.executor_request_queue.get_from_request_queue.return_value = [] + waiting_queue = Mock() + waiting_queue.__len__ = Mock(return_value=0) + + executor._fetch_and_enqueue_requests(waiting_queue, total_num_live_requests=0) + + connector.finish_prefix_loads.assert_called_once_with([17]) + waiting_queue.add_requests.assert_called_once_with([]) + if control_pending: + executor.executor_request_queue.get_from_request_queue.assert_not_called() + else: + timeout = executor.executor_request_queue.get_from_request_queue.call_args.args[0] + assert timeout.total_seconds() == 0 + + +@pytest.mark.parametrize("stop_id", [SHUTDOWN_REQUEST_ID, CONTROL_REQUEST_ID]) +def test_completion_is_applied_before_a_control_or_shutdown_boundary(stop_id: int) -> None: + executor = _executor() + executor.control_requests = [] + executor.request_accumulated = [] + executor.is_shutdown = False + completion = RequestQueueItem(PREFIX_LOAD_COMPLETION_REQUEST_ID, finished_prefix_load_ids=[17]) + + accepted = executor._handle_special_queue_items([completion, RequestQueueItem(stop_id)]) + + assert accepted == [] + assert not completion.is_normal_request + executor.kv_connector_manager.finish_prefix_loads.assert_called_once_with([17]) + + +def test_cancellation_in_completion_envelope_is_seen_before_resumption() -> None: + executor = _executor() + req = _request() + executor.active_requests = [req] + connector = executor.kv_connector_manager + order = [] + connector.defer_load_termination.side_effect = lambda request: order.append("cancel") + connector.finish_prefix_loads.side_effect = lambda ids: order.append("finish") + + executor._handle_special_queue_items( + [ + RequestQueueItem(PREFIX_LOAD_COMPLETION_REQUEST_ID, finished_prefix_load_ids=[17]), + RequestQueueItem(req.py_request_id, is_canceled_request=True), + ] + ) + + assert order == ["cancel", "finish"] + + +def test_empty_iteration_keeps_request_broadcast_fast_path() -> None: + dist = SimpleNamespace(rank=0, world_size=1) + broadcaster = RequestBroadcaster(dist, SimpleNamespace(pause=nullcontext)) + broadcaster._broadcast_requests = Mock(side_effect=AssertionError("Unexpected payload")) + + assert broadcaster.broadcast([]) == ([], None) + + +def test_completion_envelope_reaches_every_rank_without_a_compute_batch() -> None: + wire = {} + executors = [] + for rank in range(2): + executor = _executor() + executor.control_requests = [] + executor.is_shutdown = False + executor._disable_mpi = False + executor.request_accumulated = [] + executor.hang_detector = SimpleNamespace(pause=nullcontext) + + def broadcast(value, root, *, rank=rank): + assert root == 0 + if rank == 0: + wire["payload"] = deepcopy(value) + return deepcopy(wire["payload"]) + + def broadcast_count(value, root, *, rank=rank): + assert root == 0 + if rank == 0: + wire["count"] = value + return wire["count"] + + executor.dist = SimpleNamespace( + rank=rank, + world_size=2, + tp_size=2, + cp_size=1, + has_pp=False, + broadcast=broadcast, + broadcast_int64=broadcast_count, + ) + executor.request_broadcaster = RequestBroadcaster(executor.dist, executor.hang_detector) + executor.executor_request_queue = Mock() + executor.executor_request_queue.get_from_request_queue.return_value = [] + executor.kv_connector_manager.has_pending_loads.return_value = True + executor.kv_connector_manager.take_finished_prefix_loads.return_value = [17] + executors.append(executor) + + for executor in executors: + waiting_queue = Mock() + waiting_queue.__len__ = Mock(return_value=0) + executor._fetch_and_enqueue_requests(waiting_queue, total_num_live_requests=0) + executor.kv_connector_manager.finish_prefix_loads.assert_called_once_with([17]) + waiting_queue.add_requests.assert_called_once_with([]) + executors[1].kv_connector_manager.take_finished_prefix_loads.assert_not_called() diff --git a/tests/unittest/_torch/executor/test_kv_connector_reservations.py b/tests/unittest/_torch/executor/test_kv_connector_reservations.py new file mode 100644 index 000000000000..09d4419d8daa --- /dev/null +++ b/tests/unittest/_torch/executor/test_kv_connector_reservations.py @@ -0,0 +1,380 @@ +# 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. + +from types import SimpleNamespace + +import pytest + +from tensorrt_llm._torch.pyexecutor.connectors import kv_cache_connector as connector +from tensorrt_llm._torch.pyexecutor.scheduler import ScheduledRequests +from tensorrt_llm.bindings import LlmRequestState + +pytestmark = pytest.mark.cpu_only + + +class ReservationScheduler(connector.KvCacheConnectorScheduler): + def __init__(self): + super().__init__(llm_args=None) + self.protected = {} + self.releases = [] + self.queries = [] + self.is_async = True + + def reserve_prefix(self, req, num_computed_tokens, reservation_id): + self.queries.append((req.request_id, num_computed_tokens, reservation_id)) + self.protected[reservation_id] = set(range(num_computed_tokens, 96)) + return 96 - num_computed_tokens, self.is_async + + def release_prefix_reservation(self, req, reservation_id, start, end): + self.releases.append((reservation_id, start, end)) + for position in range(start, end): + self.protected[reservation_id].remove(position) + if not self.protected[reservation_id]: + del self.protected[reservation_id] + + def replace_source(self): + if self.protected: + raise RuntimeError("Source is reserved") + + def build_connector_meta(self, scheduler_output): + return scheduler_output + + def get_num_new_matched_tokens(self, req, num_computed_tokens): + return 0, False + + def update_state_after_alloc(self, req, block_ids): + pass + + def request_finished(self, req, cache_block_ids): + return False + + +class ReservationWorker(connector.KvCacheConnectorWorker): + def __init__(self): + super().__init__(llm_args=None) + self.started = [] + self.finished = [] + self.legacy_finished = ([], []) + + def register_kv_caches(self, kv_cache_tensor): + pass + + def start_load_kv(self, stream): + self.started.extend(self.get_connector_meta().prefix_loads) + + def wait_for_layer_load(self, layer_idx, stream): + pass + + def save_kv_layer(self, layer_idx, stream): + pass + + def wait_for_save(self, stream): + pass + + def get_finished(self, finished_gen_req_ids, started_loading_req_ids): + result = self.legacy_finished + self.legacy_finished = ([], []) + return result + + def get_finished_prefix_loads(self): + result = self.finished + self.finished = [] + return result + + +@pytest.fixture +def manager(monkeypatch): + monkeypatch.setattr(connector, "mpi_rank", lambda: 0) + monkeypatch.setattr(connector, "mpi_broadcast", lambda result, root: result) + monkeypatch.setattr(connector, "mpi_allgather", lambda result: [result]) + monkeypatch.setattr(connector, "mpi_world_size", lambda: 1) + result = connector.KvCacheConnectorManager(ReservationWorker(), ReservationScheduler()) + result.configure_prefix_reservations(True) + return result + + +@pytest.fixture +def req(): + return SimpleNamespace( + request_id=7, + state=LlmRequestState.CONTEXT_INIT, + cache_salt="tenant", + get_tokens=lambda beam: list(range(128)), + ) + + +def accept(manager, req): + reservation = manager.reserve_prefix(req, 32) + manager.accept_prefix_load(req, 32, 96, [[10, 11, 12], [20, -1, 22]]) + batch = ScheduledRequests() + batch.context_requests_last_chunk = [req] + manager.build_scheduler_output(batch, None) + manager.take_scheduled_requests_pending_load(batch) + assert batch.context_requests == [] + return reservation + + +def dispatch(manager): + manager.handle_metadata() + manager.mark_prefix_loads_dispatched() + manager.worker.start_load_kv(None) + + +def finish(manager): + manager.get_finished() + manager.finish_prefix_loads(manager.take_finished_prefix_loads()) + + +def test_reservation_protects_source_without_transmission(manager, req): + reservation = manager.reserve_prefix(req, 32) + assert manager.reserve_prefix(req, 32) is reservation + assert len(manager.scheduler.queries) == 1 + assert manager.worker.started == [] + assert not manager.has_pending_load(req) + with pytest.raises(RuntimeError, match="reserved"): + manager.scheduler.replace_source() + + manager.release_prefix_reservation(req) + manager.release_prefix_reservation(req) + assert manager.scheduler.releases == [(reservation.reservation_id, 32, 96)] + assert manager.pending_prefix_requests() == [] + manager.scheduler.replace_source() + + +def test_trim_releases_each_unused_range_once(manager, req): + original = manager.reserve_prefix(req, 16) + trimmed = manager.trim_prefix_reservation(req, 32, 64) + assert (trimmed.start, trimmed.end) == (32, 64) + manager.release_prefix_reservation(req) + assert manager.scheduler.releases == [ + (original.reservation_id, 16, 32), + (original.reservation_id, 64, 96), + (original.reservation_id, 32, 64), + ] + assert manager.scheduler.protected == {} + + +def test_overlapping_reservations_keep_independent_source_protection(manager, req): + other = SimpleNamespace(**vars(req)) + other.request_id = 8 + manager.reserve_prefix(req, 32) + manager.reserve_prefix(other, 32) + manager.release_prefix_reservation(req) + with pytest.raises(RuntimeError, match="reserved"): + manager.scheduler.replace_source() + manager.release_prefix_reservation(other) + manager.scheduler.replace_source() + + +def test_confirmed_async_load_survives_a_second_output_build(manager, req): + reservation = accept(manager, req) + manager.build_scheduler_output(ScheduledRequests(), None) + assert manager.worker.started == [] + assert manager.new_async_requests.loading == {} + assert manager.has_pending_loads() + dispatch(manager) + load = manager.worker.started[0] + assert load.reservation_id == reservation.reservation_id + assert (load.start, load.end, load.cache_salt) == (32, 96, "tenant") + assert load.block_ids_by_layer_group == [[10, 11, 12], [20, -1, 22]] + assert load.tokens == list(range(128)) + assert manager.worker.get_connector_meta().new_requests == [] + + +@pytest.mark.parametrize("already_bound", [False, True]) +def test_unstarted_load_can_be_released_without_dispatch(manager, req, already_bound): + reservation = accept(manager, req) + if already_bound: + manager.handle_metadata() + manager.release_unstarted_prefix_loads(req) + manager.handle_metadata() + manager.mark_prefix_loads_dispatched() + manager.worker.start_load_kv(None) + assert manager.worker.started == [] + assert manager.scheduler.releases == [(reservation.reservation_id, 32, 96)] + assert not manager.has_pending_loads() + + +def test_cancelled_load_waits_for_every_rank(manager, req): + tracker = manager._prefix_completion_tracker + tracker._size = 2 + reservation = accept(manager, req) + dispatch(manager) + assert manager.defer_load_termination(req) + manager.worker.finished = [reservation.reservation_id] + finish(manager) + manager.release_unstarted_prefix_loads(req) + assert manager.has_pending_load(req) + assert manager.scheduler.releases == [] + assert manager.take_finished_load_terminations() == [] + with pytest.raises(RuntimeError, match="owns the allocation"): + manager.reset_request_state(req) + + tracker._record(1, {reservation.reservation_id}) + completed = manager.take_finished_prefix_loads() + assert completed == [reservation.reservation_id] + assert manager.has_pending_load(req) + assert manager.scheduler.releases == [] + manager.finish_prefix_loads(completed) + assert manager.has_pending_loads() + assert req.state == LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS + assert manager.take_finished_load_terminations() == [req] + assert not manager.has_pending_loads() + assert manager.take_finished_load_terminations() == [] + assert manager.scheduler.releases == [(reservation.reservation_id, 32, 96)] + assert req.request_id not in manager.finished_async_loading_requests + + +def test_stale_completion_cannot_finish_a_replayed_allocation(manager, req): + first = accept(manager, req) + dispatch(manager) + manager.worker.finished = [first.reservation_id] + finish(manager) + assert req.state == LlmRequestState.CONTEXT_INIT + assert not manager.should_add_sequence(req) + manager.reset_request_state(req) + assert manager.should_add_sequence(req) + + second = accept(manager, req) + assert second.reservation_id > first.reservation_id + dispatch(manager) + manager.worker.finished = [first.reservation_id, first.reservation_id, 10000] + finish(manager) + assert manager.has_pending_load(req) + assert req.state == LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS + manager.worker.finished = [second.reservation_id, second.reservation_id] + finish(manager) + assert not manager.has_pending_load(req) + assert req.state == LlmRequestState.CONTEXT_INIT + assert manager.scheduler.protected == {} + + +def test_sync_load_retains_ownership_until_identity_completion(manager, req): + manager.scheduler.is_async = False + reservation = manager.reserve_prefix(req, 32) + manager.accept_prefix_load(req, 32, 96, [[1, 2, 3]]) + assert manager.scheduler_output_manager.external_loads == {req.request_id: 64} + manager.build_scheduler_output(ScheduledRequests(), None) + dispatch(manager) + assert manager.has_pending_load(req) + assert manager.defer_load_termination(req) + manager.worker.finished = [reservation.reservation_id] + finish(manager) + assert manager.take_finished_load_terminations() == [req] + assert not manager.has_pending_load(req) + + +def test_legacy_load_cancellation_retains_locally_finished_ownership(manager, req, monkeypatch): + manager.configure_prefix_reservations(False) + manager.commit_new_matched_tokens(req, 64, True) + req.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS + assert manager.defer_load_termination(req) + manager.worker.legacy_finished = ([], [req.request_id]) + remote_finished = ([], []) + monkeypatch.setattr(connector, "mpi_allgather", lambda value: [value, remote_finished]) + manager.get_finished() + assert req.request_id in manager.local_finished_async_requests.loading + assert manager.has_pending_load(req) + assert manager.take_finished_load_terminations() == [] + remote_finished = ([], [req.request_id]) + manager.get_finished() + assert manager.take_finished_load_terminations() == [req] + assert req.state == LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS + assert not manager.has_pending_loads() + + +@pytest.mark.parametrize("missing", ["reserve_prefix", "release_prefix_reservation", "worker"]) +def test_partial_reservation_capability_is_rejected(manager, monkeypatch, missing): + manager._prefix_capability = None + if missing == "worker": + monkeypatch.setattr( + ReservationWorker, + "get_finished_prefix_loads", + connector.KvCacheConnectorWorker.get_finished_prefix_loads, + ) + else: + monkeypatch.setattr( + ReservationScheduler, missing, getattr(connector.KvCacheConnectorScheduler, missing) + ) + with pytest.raises(ValueError, match="prefix reservations require"): + manager.configure_prefix_reservations(True) + + +def test_legacy_connectors_do_not_enable_reservations(manager, monkeypatch): + manager._prefix_capability = None + for name in ("reserve_prefix", "release_prefix_reservation"): + monkeypatch.setattr( + ReservationScheduler, name, getattr(connector.KvCacheConnectorScheduler, name) + ) + monkeypatch.setattr( + ReservationWorker, + "get_finished_prefix_loads", + connector.KvCacheConnectorWorker.get_finished_prefix_loads, + ) + manager.configure_prefix_reservations(True) + assert not manager.prefix_reservations_enabled + + +def test_reserving_and_releasing_do_not_gather_worker_state(manager, req, monkeypatch): + def unexpected_collective(*args, **kwargs): + pytest.fail("Reservation validation or release entered a collective") + + monkeypatch.setattr(connector, "mpi_allgather", unexpected_collective) + reservation = manager.reserve_prefix(req, 32) + monkeypatch.setattr(connector, "mpi_broadcast", unexpected_collective) + manager.trim_prefix_reservation(req, 32, 64) + manager.release_prefix_reservation(req) + assert manager.scheduler.releases == [ + (reservation.reservation_id, 64, 96), + (reservation.reservation_id, 32, 64), + ] + + +@pytest.mark.parametrize("with_load", [False, True]) +def test_prefix_polling_adds_no_collective(manager, req, monkeypatch, with_load): + if with_load: + reservation = accept(manager, req) + dispatch(manager) + manager.worker.finished = [reservation.reservation_id] + calls = [] + + def gather(value): + calls.append(value) + return [value] + + monkeypatch.setattr(connector, "mpi_allgather", gather) + manager.get_finished() + assert calls == [([], [])] + if with_load: + assert req.state == LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS + manager.finish_prefix_loads(manager.take_finished_prefix_loads()) + assert req.state == LlmRequestState.CONTEXT_INIT + + +def test_retirement_requires_local_transfer_completion(manager, req): + reservation = accept(manager, req) + dispatch(manager) + with pytest.raises(RuntimeError, match="has not finished locally"): + manager.finish_prefix_loads([reservation.reservation_id]) + assert manager.has_pending_load(req) + assert manager.scheduler.releases == [] + + +@pytest.mark.parametrize("answer", [(-1, False), (True, False), (0, True), (32, 1)]) +def test_invalid_reservation_answer_is_rejected(manager, req, monkeypatch, answer): + monkeypatch.setattr(manager.scheduler, "reserve_prefix", lambda *args: answer) + with pytest.raises(ValueError): + manager.reserve_prefix(req, 32) + assert manager.get_prefix_reservation(req) is None diff --git a/tests/unittest/_torch/executor/test_kv_connector_v2_prefix.py b/tests/unittest/_torch/executor/test_kv_connector_v2_prefix.py index 89d7b991f74e..6986d77f4034 100644 --- a/tests/unittest/_torch/executor/test_kv_connector_v2_prefix.py +++ b/tests/unittest/_torch/executor/test_kv_connector_v2_prefix.py @@ -12,28 +12,10 @@ # 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. -"""Unit tests for the KV connector prefix. - -The connector is asked in ``prepare_resources``, on the batch the forward pass -will run, which is downstream of every stage that can drop a request. The -*asked => scheduled => eventually request_finished* invariant therefore holds, -an offer is never abandoned, and no ``cancel_load`` is needed. - -Two properties of the allocation are what these tests pin. - -* **Pages are allocated per context chunk**, deliberately -- that is what - chunked prefill is for -- so an offer reaching beyond the chunk needs a - bounded grow, and the grow can fail. -* **The local match is token-granular**: ``num_committed_tokens`` is not - floored to whole shared blocks, so the arithmetic can go negative. - -``FakeRequest`` reproduces ``LlmRequest``'s chunk arithmetic including -``setContextChunkSize``'s non-negative check and ``setPrepopulatedPromptLen``'s -block-alignment assertion, so a version of this code that violates either fails -here rather than only on hardware. -""" +"""Connector prefix reservation, allocation and legacy final-batch query tests.""" from types import SimpleNamespace +from unittest.mock import Mock import pytest @@ -60,6 +42,13 @@ def __init__(self, committed=0, capacity=None, grow_ok=True): self.is_active = True self.grow_ok = grow_ok self.resize_calls = [] + self.closed = False + + def close(self) -> None: + self.closed = True + + def discard_pending_stats(self) -> None: + pass def resume(self, cuda_stream): self.is_active = True @@ -159,6 +148,13 @@ class FakeConnectorManager: """Records the calls the prefix path makes, in order.""" def __init__(self, num_matched=0, load_async=False, add_sequence=True): + self.prefix_reservations_enabled = False + self.reservations = {} + self.reservation_requests = {} + self.next_reservation_id = 1 + self.releases = [] + self.accepted = [] + self.dispatched = set() self.num_matched = num_matched self.load_async = load_async self.add_sequence = add_sequence @@ -168,6 +164,65 @@ def __init__(self, num_matched=0, load_async=False, add_sequence=True): self.alloc_by_group = [] self.forgotten = [] + def reserve_prefix(self, request: FakeRequest, local_end: int) -> SimpleNamespace | None: + if request.request_id in self.reservations: + return self.reservations[request.request_id] + self.queries.append((request.request_id, local_end)) + if not self.num_matched: + return None + reservation = SimpleNamespace( + reservation_id=self.next_reservation_id, + request_id=request.request_id, + start=local_end, + end=local_end + self.num_matched, + is_async=self.load_async, + ) + self.next_reservation_id += 1 + self.reservations[request.request_id] = reservation + self.reservation_requests[request.request_id] = request + return reservation + + def get_prefix_reservation(self, request: FakeRequest) -> SimpleNamespace | None: + return self.reservations.get(request.request_id) + + def trim_prefix_reservation( + self, request: FakeRequest, start: int, end: int + ) -> SimpleNamespace: + reservation = self.reservations[request.request_id] + if reservation.start < start: + self.releases.append((reservation.reservation_id, reservation.start, start)) + if end < reservation.end: + self.releases.append((reservation.reservation_id, end, reservation.end)) + reservation.start, reservation.end = start, end + return reservation + + def release_prefix_reservation(self, request: FakeRequest) -> None: + reservation = self.reservations.pop(request.request_id, None) + self.reservation_requests.pop(request.request_id, None) + if reservation is not None: + self.releases.append((reservation.reservation_id, reservation.start, reservation.end)) + + def pending_prefix_requests(self) -> list[FakeRequest]: + return list(self.reservation_requests.values()) + + def accept_prefix_load( + self, request: FakeRequest, start: int, end: int, block_ids_by_layer_group: list[list[int]] + ) -> None: + reservation = self.reservations.pop(request.request_id) + self.reservation_requests.pop(request.request_id) + self.accepted.append((reservation, block_ids_by_layer_group)) + request.py_num_connector_matched_tokens = end - start + + def release_unstarted_prefix_loads(self, request: FakeRequest) -> None: + self.accepted = [ + entry + for entry in self.accepted + if entry[0].request_id != request.request_id or request.request_id in self.dispatched + ] + + def has_pending_load(self, request: FakeRequest) -> bool: + return any(entry[0].request_id == request.request_id for entry in self.accepted) + def query_num_new_matched_tokens(self, request, num_computed_tokens): self.queries.append((request.request_id, num_computed_tokens)) return self.num_matched, self.load_async @@ -207,6 +262,16 @@ def make_manager(connector, num_extra_kv_tokens=0, is_draft=False): manager.kv_cache_map = {} manager.enable_block_reuse = True manager.conversation_manager = None + manager._has_cp_helix = False + manager._allocated_draft_lens = {} + manager._request_stats_enabled_ids = set() + manager._fresh_pages_filled = {} + manager._disagg_receive_ready = {} + manager._early_freed_index_requests = set() + manager.impl = Mock() + manager.index_mapper = Mock() + manager._fill_fresh_kv_pages = Mock() + manager._log_window_crossing = Mock() manager._stream = SimpleNamespace(cuda_stream=0) # One layer group, the shape every non-VSWA, non-hybrid model has. The real # accessor reads `impl.layer_grouping`, which only a pool allocation fills @@ -795,3 +860,163 @@ def test_scratch_reuse_survives_without_a_connector(self): assert manager.prepare_context(req) assert kv_cache.enable_swa_scratch_reuse is True + + +class TestPrefixReservations: + @staticmethod + def prepare( + num_matched: int = 64, + committed: int = 0, + capacity: int = 0, + grow_ok: bool = True, + load_async: bool = False, + ) -> tuple[KVCacheManagerV2, FakeConnectorManager, FakeRequest, FakeKvCache]: + connector = FakeConnectorManager(num_matched=num_matched, load_async=load_async) + connector.prefix_reservations_enabled = True + manager = make_manager(connector) + req = FakeRequest() + cache = FakeKvCache(committed=committed, capacity=capacity, grow_ok=grow_ok) + manager.kv_cache_map[req.request_id] = cache + assert manager.prepare_context(req) + return manager, connector, req, cache + + @pytest.mark.parametrize("load_async", [False, True]) + def test_reserve_before_budget_and_accept_after_final_trim(self, load_async: bool) -> None: + manager, connector, req, cache = self.prepare(load_async=load_async) + assert req.context_current_position == 64 + assert req.py_connector_served_position == 0 + assert cache.history_length == 0 + assert connector.accepted == [] + assert connector.allocs == [] + + req.context_chunk_size = 64 + assert manager.resize_context(req, 64) + assert (cache.capacity, cache.history_length) == (128, 64) + batch = scheduled(req) + manager._run_kv_connector_hooks(batch) + assert connector.accepted == [] + req.context_chunk_size = 32 + manager.report_batch_to_connector(batch) + + assert req.py_connector_served_position == 64 + assert req.context_chunk_size == 32 + assert len(connector.allocs) == 1 + reservation, groups = connector.accepted[0] + assert (reservation.start, reservation.end, reservation.is_async) == (0, 64, load_async) + assert groups == [[]] + assert connector.reservations == {} + manager.report_batch_to_connector(batch) + assert len(connector.accepted) == 1 + assert len(connector.allocs) == 1 + + @pytest.mark.parametrize("capacity", [0, PROMPT_LEN]) + def test_rejection_rewinds_even_without_capacity_growth(self, capacity: int) -> None: + manager, connector, req, cache = self.prepare(capacity=capacity) + assert manager.resize_context(req, 32) + assert cache.history_length == 64 + manager.release_unused_connector_reservations(set()) + assert connector.releases == [(1, 0, 64)] + assert connector.accepted == [] + assert req.request_id not in manager.kv_cache_map + assert cache.closed + assert req.context_current_position == 0 + assert req.prepopulated_prompt_len == 0 + assert req.context_chunk_size == PROMPT_LEN + assert req.py_connector_served_position == 0 + manager.release_unused_connector_reservations(set()) + assert connector.releases == [(1, 0, 64)] + + def test_failed_allocation_releases_credit(self) -> None: + manager, connector, req, cache = self.prepare(grow_ok=False) + assert not manager.resize_context(req, 32) + assert connector.releases == [(1, 0, 64)] + assert cache.closed + assert req.context_current_position == 0 + assert req.request_id not in manager.kv_cache_map + + def test_revert_discards_unfilled_history_without_capacity_growth(self) -> None: + manager, connector, req, cache = self.prepare(capacity=PROMPT_LEN) + assert manager.resize_context(req, 32) + assert req.py_ctx_pre_resize_cap is None + assert not manager.revert_allocate_context(req) + assert cache.closed + assert connector.releases == [(1, 0, 64)] + assert req.context_current_position == 0 + + @pytest.mark.parametrize( + "committed,offered,expected_end", + [(17, 47, 64), (17, 14, 17), (0, PROMPT_LEN, PROMPT_LEN - TOKENS_PER_BLOCK)], + ) + def test_partial_local_prefix_and_last_prompt_token( + self, committed: int, offered: int, expected_end: int + ) -> None: + manager, connector, req, cache = self.prepare( + num_matched=offered, committed=committed, capacity=PROMPT_LEN + ) + assert req.context_current_position == expected_end + assert req.context_remaining_length == PROMPT_LEN - expected_end + assert req.py_connector_served_position == 0 + reservation = connector.get_prefix_reservation(req) + if expected_end == committed: + assert reservation is None + assert connector.releases == [(1, committed, committed + offered)] + else: + assert (reservation.start, reservation.end) == (committed, expected_end) + + def test_retry_in_same_attempt_queries_once(self) -> None: + manager, connector, req, _ = self.prepare() + assert manager.prepare_context(req) + assert req.context_current_position == 64 + assert connector.queries == [(req.request_id, 0)] + + def test_dispatched_load_cannot_lose_its_allocation(self) -> None: + manager, connector, req, cache = self.prepare(load_async=True) + assert manager.resize_context(req, 32) + req.context_chunk_size = 32 + manager.report_batch_to_connector(scheduled(req)) + connector.dispatched.add(req.request_id) + manager.release_unused_connector_reservations(set()) + with pytest.raises(RuntimeError, match="while request .* is loading"): + manager.free_resources(req) + assert not cache.closed + assert req.py_connector_served_position == 64 + assert manager.kv_cache_map[req.request_id] is cache + + def test_final_batch_releases_only_rejected_requests(self) -> None: + manager, connector, req, cache = self.prepare() + assert manager.resize_context(req, 32) + req.context_chunk_size = 32 + other = FakeRequest(request_id=1) + other_cache = FakeKvCache() + manager.kv_cache_map[other.request_id] = other_cache + assert manager.prepare_context(other) + assert manager.resize_context(other, 32) + manager.report_batch_to_connector(scheduled(req)) + assert len(connector.accepted) == 1 + assert connector.releases == [(2, 0, 64)] + assert not cache.closed + assert other_cache.closed + assert other.context_current_position == 0 + + +def test_disagg_metadata_build_keeps_pending_context_reservation() -> None: + manager, connector, req, cache = TestPrefixReservations.prepare() + assert manager.resize_context(req, 32) + manager.report_batch_to_connector(scheduled(), finalize_prefix_reservations=False) + assert connector.get_prefix_reservation(req) is not None + assert connector.releases == [] + assert not cache.closed + + +def test_final_admission_rejects_an_invalid_local_allocation() -> None: + manager, connector, req, cache = TestPrefixReservations.prepare() + assert manager.resize_context(req, 32) + req.context_chunk_size = 32 + cache.history_length = 0 + with pytest.raises(RuntimeError, match="no allocation for its reserved KV prefix"): + manager.report_batch_to_connector(scheduled(req)) + assert connector.accepted == [] + assert connector.get_prefix_reservation(req) is not None + assert req.py_connector_served_position == 0 + manager.release_unused_connector_reservations(set()) + assert connector.releases == [(1, 0, 64)] diff --git a/tests/unittest/_torch/executor/test_kv_connector_v2_prefix_real_manager.py b/tests/unittest/_torch/executor/test_kv_connector_v2_prefix_real_manager.py index db2f83697e23..a6ea3276d7ef 100644 --- a/tests/unittest/_torch/executor/test_kv_connector_v2_prefix_real_manager.py +++ b/tests/unittest/_torch/executor/test_kv_connector_v2_prefix_real_manager.py @@ -20,6 +20,7 @@ """ import gc +from types import SimpleNamespace import pytest import torch @@ -53,6 +54,13 @@ class FakeConnectorManager: """Records what the prefix path tells the connector, in order.""" def __init__(self, num_matched=OFFER_TOKENS, load_async=False): + self.prefix_reservations_enabled = False + self.reservations = {} + self.reservation_requests = {} + self.next_reservation_id = 1 + self.releases = [] + self.accepted = [] + self.dispatched = set() self.num_matched = num_matched self.load_async = load_async self.queries = [] @@ -61,6 +69,63 @@ def __init__(self, num_matched=OFFER_TOKENS, load_async=False): self.allocs_by_group = [] self.forgotten = [] + def reserve_prefix(self, request: LlmRequest, local_end: int) -> SimpleNamespace | None: + if request.request_id in self.reservations: + return self.reservations[request.request_id] + self.queries.append((request.request_id, local_end)) + if not self.num_matched: + return None + reservation = SimpleNamespace( + reservation_id=self.next_reservation_id, + request_id=request.request_id, + start=local_end, + end=local_end + self.num_matched, + is_async=self.load_async, + ) + self.next_reservation_id += 1 + self.reservations[request.request_id] = reservation + self.reservation_requests[request.request_id] = request + return reservation + + def get_prefix_reservation(self, request: LlmRequest) -> SimpleNamespace | None: + return self.reservations.get(request.request_id) + + def trim_prefix_reservation(self, request: LlmRequest, start: int, end: int) -> SimpleNamespace: + reservation = self.reservations[request.request_id] + if reservation.start < start: + self.releases.append((reservation.reservation_id, reservation.start, start)) + if end < reservation.end: + self.releases.append((reservation.reservation_id, end, reservation.end)) + reservation.start, reservation.end = start, end + return reservation + + def release_prefix_reservation(self, request: LlmRequest) -> None: + reservation = self.reservations.pop(request.request_id, None) + self.reservation_requests.pop(request.request_id, None) + if reservation is not None: + self.releases.append((reservation.reservation_id, reservation.start, reservation.end)) + + def pending_prefix_requests(self) -> list[LlmRequest]: + return list(self.reservation_requests.values()) + + def accept_prefix_load( + self, request: LlmRequest, start: int, end: int, block_ids_by_layer_group: list[list[int]] + ) -> None: + reservation = self.reservations.pop(request.request_id) + self.reservation_requests.pop(request.request_id) + self.accepted.append((reservation, block_ids_by_layer_group)) + request.py_num_connector_matched_tokens = end - start + + def release_unstarted_prefix_loads(self, request: LlmRequest) -> None: + self.accepted = [ + entry + for entry in self.accepted + if entry[0].request_id != request.request_id or request.request_id in self.dispatched + ] + + def has_pending_load(self, request: LlmRequest) -> bool: + return any(entry[0].request_id == request.request_id for entry in self.accepted) + def query_num_new_matched_tokens(self, request, num_computed_tokens): self.queries.append((request.py_request_id, num_computed_tokens)) return self.num_matched, self.load_async @@ -634,3 +699,167 @@ def test_a_served_prefix_survives_re_entry_under_a_sliding_window(): del mgr gc.collect() torch.cuda.empty_cache() + + +@pytest.mark.parametrize("preallocated", [False, True]) +def test_rejected_connector_candidate_releases_and_rewinds( + manager: KVCacheManagerV2, connector: FakeConnectorManager, preallocated: bool +) -> None: + request = make_request() + if preallocated: + assert schedule(manager, request) + connector.prefix_reservations_enabled = True + assert schedule(manager, request) + cache = manager.kv_cache_map[request.request_id] + assert cache.history_length == OFFER_TOKENS + if preallocated: + assert request.py_ctx_pre_resize_cap is None + assert connector.accepted == [] + + manager.release_unused_connector_reservations(set()) + + assert request.request_id not in manager.kv_cache_map + assert connector.releases == [(1, 0, OFFER_TOKENS)] + assert request.context_current_position == 0 + assert request.prepopulated_prompt_len == 0 + assert request.context_chunk_size == PROMPT_LEN + assert request.py_connector_served_position == 0 + assert schedule(manager, request) + assert connector.get_prefix_reservation(request).reservation_id == 2 + assert connector.queries == [(request.request_id, 0), (request.request_id, 0)] + + +def test_token_budget_rejection_releases_real_cache( + manager: KVCacheManagerV2, connector: FakeConnectorManager +) -> None: + from tensorrt_llm._torch.pyexecutor.scheduler.scheduler_v2 import KVCacheV2Scheduler + from tensorrt_llm.llmapi.llm_args import CapacitySchedulerPolicy + + connector.prefix_reservations_enabled = True + request = make_request() + scheduler = KVCacheV2Scheduler( + max_batch_size=4, + max_num_tokens=32, + kv_cache_manager=manager, + scheduler_policy=CapacitySchedulerPolicy.MAX_UTILIZATION, + ) + + output = scheduler.schedule_request([request], set()) + + assert output.context_requests == [] + assert request.request_id not in manager.kv_cache_map + assert connector.releases == [(1, 0, OFFER_TOKENS)] + assert request.context_current_position == 0 + assert connector.accepted == [] + + +def test_reserved_prefix_tracks_range_and_allocation( + manager: KVCacheManagerV2, connector: FakeConnectorManager +) -> None: + from tensorrt_llm._torch.pyexecutor.llm_request import rewind_context_after_cache_drop + + connector.prefix_reservations_enabled = True + connector.num_matched = PROMPT_LEN + request = make_request() + assert schedule(manager, request) + batch = run(manager, request) + assert connector.accepted == [] + manager.report_batch_to_connector(batch) + first, first_groups = connector.accepted[0] + assert (first.start, first.end) == (0, 64) + assert connector.releases == [(first.reservation_id, 64, PROMPT_LEN)] + assert all(slot >= 0 for _, slot in valid_page_slots(first_groups[0])) + assert len(first_groups[0]) == 3 + + manager.free_resources(request) + rewind_context_after_cache_drop(request, TOKENS_PER_BLOCK) + assert schedule(manager, request) + manager.report_batch_to_connector(run(manager, request)) + + replay, replay_groups = connector.accepted[0] + assert replay.reservation_id != first.reservation_id + assert (replay.start, replay.end) == (first.start, first.end) + assert len(replay_groups[0]) == 3 + assert connector.queries == [(request.request_id, 0), (request.request_id, 0)] + assert len(connector.allocs) == 2 + + +def test_dispatched_load_retains_real_destination_pages( + manager: KVCacheManagerV2, connector: FakeConnectorManager +) -> None: + connector.prefix_reservations_enabled = True + connector.load_async = True + request = make_request() + assert schedule(manager, request) + manager.report_batch_to_connector(run(manager, request)) + connector.dispatched.add(request.request_id) + pages_before = manager.get_page_indices_by_layer_group(request) + + with pytest.raises(RuntimeError, match="while request .* is loading"): + manager.free_resources(request) + + other = make_request(request_id=2) + assert schedule(manager, other) + pages_after = manager.get_page_indices_by_layer_group(request) + other_pages = manager.get_page_indices_by_layer_group(other) + assert pages_after == pages_before + for held, allocated in zip(pages_before, other_pages): + assert {slot for _, slot in valid_page_slots(held)}.isdisjoint( + slot for _, slot in valid_page_slots(allocated) + ) + assert manager.kv_cache_map[request.request_id].is_active + + +def test_reserved_vswa_prefix_reports_live_group_ordinals( + vswa_manager: KVCacheManagerV2, vswa_connector: FakeConnectorManager +) -> None: + vswa_connector.prefix_reservations_enabled = True + request = make_request(prompt_len=VSWA_PROMPT_LEN) + assert schedule(vswa_manager, request) + vswa_manager.report_batch_to_connector(run(vswa_manager, request)) + reservation, groups = vswa_connector.accepted[0] + sliding, full = _sliding_and_full(vswa_manager) + assert (reservation.start, reservation.end) == (0, VSWA_OFFER) + assert len(groups[sliding]) == len(groups[full]) == VSWA_PROMPT_LEN // TOKENS_PER_BLOCK + stale_end = (VSWA_OFFER + 1 - VSWA_WINDOW) // TOKENS_PER_BLOCK + assert groups[sliding][:stale_end] == [BAD_PAGE_INDEX] * stale_end + assert all(slot >= 0 for slot in groups[sliding][stale_end:]) + assert all(slot >= 0 for slot in groups[full]) + + +def test_real_kv_pressure_rejects_reserved_prefix_without_transmission() -> None: + connector = FakeConnectorManager(num_matched=64) + connector.prefix_reservations_enabled = True + manager = make_manager( + connector, + kv_cache_config=KvCacheConfig(max_tokens=64, enable_block_reuse=True), + ) + blocker = None + try: + available_blocks = manager.get_num_free_blocks() + assert available_blocks > 0 + request = make_request() + assert manager.prepare_context(request) + assert request.context_current_position == 64 + + # Preparation can already hold pages. Fill the remaining pool until + # the allocator refuses another block, preserving those real holds. + blocker = manager.impl.create_kv_cache() + assert blocker.resume(manager._stream.cuda_stream) + for num_blocks in range(1, available_blocks + 2): + if not blocker.resize(num_blocks * TOKENS_PER_BLOCK): + break + else: + pytest.fail("The blocker did not exhaust the resolved KV pool") + assert blocker.capacity > 0 + + assert not manager.resize_context(request, 32) + assert blocker.is_active + assert connector.accepted == [] + assert connector.releases == [(1, 0, 64)] + assert request.request_id not in manager.kv_cache_map + assert request.context_current_position == 0 + finally: + if blocker is not None: + blocker.close() + manager.shutdown() diff --git a/tests/unittest/_torch/executor/test_prefix_load_completion.py b/tests/unittest/_torch/executor/test_prefix_load_completion.py new file mode 100644 index 000000000000..ebee8cd46a61 --- /dev/null +++ b/tests/unittest/_torch/executor/test_prefix_load_completion.py @@ -0,0 +1,176 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from collections import defaultdict, deque + +import pytest + +from tensorrt_llm._torch.pyexecutor.connectors.prefix_load_completion import ( + PrefixLoadCompletionTracker, +) + +pytestmark = pytest.mark.cpu_only + + +class Report: + def __init__(self, ids): + self.ids = ids + self.received = False + self.ready = True + + def test(self): + return self.received, None + + def wait(self): + assert self.received + + def irecv(self): + return ReceivedReport(self) + + +class ReceivedReport: + def __init__(self, report): + self.report = report + + def test(self): + if not self.report.ready: + return False, None + self.report.received = True + return True, self.report.ids + + def wait(self): + assert self.report.ready + + +class Mailbox: + def __init__(self, rank, queues, size=2): + self.rank = rank + self.queues = queues + self.size = size + self.probes = 0 + self.sent = [] + self.freed = False + + def Get_rank(self): + return self.rank + + def Get_size(self): + return self.size + + def isend(self, ids, dest, tag): + assert dest == 0 and tag == 0 + report = Report(ids) + self.queues[self.rank].append(report) + self.sent.append(report) + return report + + def improbe(self, source, tag): + assert tag == 0 + self.probes += 1 + return self.queues[source].popleft() if self.queues[source] else None + + def Free(self): + self.freed = True + + +@pytest.fixture +def trackers(): + queues = defaultdict(deque) + return [PrefixLoadCompletionTracker(Mailbox(rank, queues)) for rank in range(2)] + + +def test_slow_worker_does_not_block_other_completed_loads(trackers): + leader, worker = trackers + for tracker in trackers: + tracker.track(1) + tracker.track(2) + leader.report({1, 2}) + worker.report({2}) + worker.poll() + assert leader.take_completed() == [2] + leader.forget(2) + assert leader.take_completed() == [] + worker.report({1}) + worker.poll() + assert leader.take_completed() == [1] + + +def test_pending_send_preserves_new_reports_without_waiting(trackers): + leader, worker = trackers + for reservation_id in (1, 2): + leader.track(reservation_id) + leader.report({1, 2}) + worker.report({1}) + worker.poll() + worker.report({2}) + worker.poll() + assert len(worker._comm.sent) == 1 + assert worker._comm.sent[0].ids == {1} + assert leader.take_completed() == [1] + worker.poll() + assert leader.take_completed() == [2] + + +def test_matched_receive_is_polled_without_waiting(trackers): + leader, worker = trackers + leader.track(1) + leader.report({1}) + worker.report({1}) + worker.poll() + report = worker._comm.sent[0] + report.ready = False + assert leader.take_completed() == [] + assert leader.take_completed() == [] + report.ready = True + assert leader.take_completed() == [1] + + +def test_stale_or_duplicate_reports_cannot_retire_another_load(trackers): + leader, worker = trackers + leader.track(2) + leader.report({2}) + worker.report({1, 1000}) + worker.poll() + assert leader.take_completed() == [] + worker.report({2}) + worker.poll() + assert leader.take_completed() == [2] + leader.forget(2) + worker.report({2}) + worker.poll() + leader.track(3) + leader.report({3}) + assert leader.take_completed() == [] + + +def test_idle_poll_does_not_send_or_probe(trackers): + for tracker in trackers: + tracker.poll() + assert tracker.take_completed() == [] + assert tracker._comm.probes == 0 + assert tracker._comm.sent == [] + + +def test_single_worker_needs_no_communicator(): + tracker = PrefixLoadCompletionTracker() + tracker.track(4) + assert tracker.take_completed() == [] + tracker.report({4}) + assert tracker.take_completed() == [4] + tracker.forget(4) + tracker.report({4}) + assert tracker.take_completed() == [] + tracker.close() + + +def test_shutdown_completes_delivered_sends_and_frees_communicator(trackers): + leader, worker = trackers + leader.track(1) + leader.report({1}) + worker.report({1}) + worker.poll() + assert leader.take_completed() == [1] + for tracker in trackers: + comm = tracker._comm + tracker.close() + assert comm.freed