From 7c87a01a5921a5b6afa8e489b0e05d24db5ebe04 Mon Sep 17 00:00:00 2001 From: Athena Cai Date: Fri, 7 Aug 2026 15:59:40 +0000 Subject: [PATCH 1/3] [TRTLLM-12499][feat] Pipelined KVCache transfer for disaggregated serving in Python Cache Transceiver Instead of waiting for all prefill chunks to complete before starting KV cache transfer, each chunk's KV data is transferred to the generation server immediately after its prefill completes. This overlaps GPU compute with RDMA transfer, hiding transfer latency behind prefill computation. Only the last chunk's transfer remains on the critical path. The feature is gated behind `enable_pipelined_transfer` on `CacheTransceiverConfig` and is implemented in `KvCacheTransceiverV2` only. It requires `schedule_style: generation_first`, `enable_chunked_prefill: true`, `beam_width == 1`, the NIXL backend, `kv_cache_bounce_size_mb == 0`, `pipeline_parallel_size == 1` on the sender, and a non-Mamba/hybrid cache manager. Each requirement is enforced at startup or per request. Squashed from 15 commits: - Chunking is sender-side only; the generation server posts a single receive covering the whole prompt and completes on `is_last_slice`. - `KVSlice` now describes one chunk rather than one whole request, gaining `total_blocks` and a meaningful `is_last_slice`. `prompt_len` became required on the session args so SWA can compute the stale-block boundary. - `project_blocks_to_global_chunk` intersects ranges instead of indexing, so resident-suffix block lists (sliding window groups, prefix reuse, incremental allocation) project correctly onto a global chunk. - The first slice always extends back to block 0, so a context-side prefix-reuse hit does not leave `[0, prepopulated_prompt_len)` unsent. - Source blocks are capped at the computed chunk boundary before SWA trimming, normalizing V1's full-prompt reservation against V2's incremental allocation. - `KV_AGENT_RESULT` carries `sender_slice_id` and `receiver_slice_id` separately, making per-chunk RDMA failures attributable. Behavior-neutral for the monolithic receiver. - KV transfer activity is modeled by transceiver session membership rather than `LlmRequestState`, so mid-prefill cancellation and transfer-timeout monitoring work during the pipelined phase. - A retired send session cannot be silently re-created, since closing it drops the peer's `RecvReqInfo` and the receiver never re-registers. - `TxSession.dispatch_lock` serializes chunk dispatch across the executor thread and the late-peer replay path, so a newer slice cannot reach a peer's queue ahead of an older one. - Transceiver configuration resolution happens early and idempotently, and backend/runtime compatibility validation is centralized. Signed-off-by: Athena Cai --- .../_torch/auto_deploy/shim/ad_executor.py | 1 + .../_torch/disaggregation/base/transfer.py | 67 +- .../_torch/disaggregation/native/transfer.py | 359 +++-- .../_torch/disaggregation/transceiver.py | 216 ++- tensorrt_llm/_torch/pyexecutor/_util.py | 9 +- .../_torch/pyexecutor/kv_cache_transceiver.py | 122 +- tensorrt_llm/_torch/pyexecutor/llm_request.py | 4 + tensorrt_llm/_torch/pyexecutor/py_executor.py | 112 +- .../_torch/pyexecutor/py_executor_creator.py | 4 + tensorrt_llm/commands/serve.py | 5 +- tensorrt_llm/llmapi/disagg_utils.py | 47 +- tensorrt_llm/llmapi/llm_args.py | 8 + .../usage/llm_args_golden_manifest.json | 7 + .../accuracy/test_disaggregated_serving.py | 95 ++ .../test_lists/test-db/l0_dgx_b200.yml | 6 + .../test_disagg_index_mapper_early_release.py | 2 + .../test_disagg_inflight_cancel_gate.py | 2 + tests/unittest/disaggregated/test_bounce.py | 29 +- .../disaggregated/test_chunked_transfer.py | 1403 +++++++++++++++++ .../disaggregated/test_disagg_utils.py | 93 ++ .../disaggregated/test_kv_transfer.py | 639 +++++++- .../test_transceiver_bounded_polling.py | 1 + 22 files changed, 3016 insertions(+), 215 deletions(-) create mode 100644 tests/unittest/disaggregated/test_chunked_transfer.py diff --git a/tensorrt_llm/_torch/auto_deploy/shim/ad_executor.py b/tensorrt_llm/_torch/auto_deploy/shim/ad_executor.py index 1ed345a1914f..8fd2e95b11b1 100644 --- a/tensorrt_llm/_torch/auto_deploy/shim/ad_executor.py +++ b/tensorrt_llm/_torch/auto_deploy/shim/ad_executor.py @@ -1281,6 +1281,7 @@ def create_autodeploy_executor( attention_type_cpp, cache_transceiver_config, mamba_cache_manager=None, + enable_chunked_prefill=getattr(ad_config, "enable_chunked_prefill", False), ) # Guided (structured) decoding. diff --git a/tensorrt_llm/_torch/disaggregation/base/transfer.py b/tensorrt_llm/_torch/disaggregation/base/transfer.py index 320f7e067bdf..cecb3a443086 100644 --- a/tensorrt_llm/_torch/disaggregation/base/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/base/transfer.py @@ -11,6 +11,36 @@ from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest +def project_blocks_to_global_chunk( + block_ids: np.ndarray, + chunk_block_offset: int, + chunk_block_count: int, + resident_block_end: int, +) -> np.ndarray: + """Project a global block chunk into a suffix-resident block list. + + ``block_ids`` represents the resident suffix of the logical range + ``[0, resident_block_end)``. ``chunk_block_offset`` and + ``chunk_block_count`` describe a chunk in that global coordinate space. + """ + if chunk_block_count <= 0 or len(block_ids) == 0: + return block_ids[:0] + + resident_start = max(0, resident_block_end - len(block_ids)) + resident_end = resident_block_end + chunk_start = chunk_block_offset + chunk_end = chunk_start + chunk_block_count + + overlap_start = max(chunk_start, resident_start) + overlap_end = min(chunk_end, resident_end) + if overlap_start >= overlap_end: + return block_ids[:0] + + local_start = overlap_start - resident_start + local_end = overlap_end - resident_start + return block_ids[local_start:local_end] + + @dataclass class TokenRange: """Range of tokens in the sequence dimension.""" @@ -25,6 +55,25 @@ def __post_init__(self): raise ValueError(f"Invalid range: [{self.start}, {self.end})") +def derive_chunk_block_coords( + token_range: Optional[TokenRange], + tokens_per_block: int, +) -> tuple[int, int]: + """Derive global chunk block offset and count from a block-aligned token_range.""" + if token_range is None: + return 0, 0 + if tokens_per_block <= 0: + raise ValueError("tokens_per_block must be positive") + if token_range.start % tokens_per_block != 0 or token_range.end % tokens_per_block != 0: + raise ValueError( + f"token_range [{token_range.start}, {token_range.end}) must be " + f"block-aligned with tokens_per_block={tokens_per_block}" + ) + chunk_offset = token_range.start // tokens_per_block + chunk_block_count = (token_range.end - token_range.start) // tokens_per_block + return chunk_offset, chunk_block_count + + @dataclass class LayerRange: """Range of layers to transfer.""" @@ -44,7 +93,9 @@ class KVSlice: """A KV cache slice covering token_range = [start, end) of one request. Single-slice transfer uses [0, prompt_len) with is_last_slice=True; - multi-slice transfers split token_range and mark the last slice. + multi-slice (pipelined) transfers split token_range and mark the last slice. + For pipelined chunks, token_range is block-aligned and encodes the global + chunk position; derive block offset/count via derive_chunk_block_coords(). Per-layer token starts are NOT encoded in token_range — they are derived from block count by the sender: @@ -65,6 +116,7 @@ class KVSlice: ) # Physical block IDs per layer group, each np.ndarray(dtype=np.int64) is_last_slice: bool = False mamba_state_index: Optional[int] = None + total_blocks: Optional[int] = None class SessionStatus(Enum): @@ -104,7 +156,7 @@ class SessionArgsBase: params: DisaggregatedParams # Captured from LlmRequest.prompt_len; needed for SWA stale_end derivation. - prompt_len: Optional[int] = None + prompt_len: int beam_width: int = 1 @@ -158,7 +210,16 @@ def __init__(self, sender: SenderBase, args: SessionArgsBase): self._sender = sender @abstractmethod - def send(self, slice: KVSlice) -> None: ... + def send(self, slice: KVSlice) -> None: + """Send a KV slice. + + Args: + slice: The KV slice describing which source blocks to send. + For pipelined chunks, ``token_range`` is the shared sender-side + chunk cursor; each layer group projects it into its own + resident/windowed source and destination block ranges. + """ + ... @abstractmethod def wait_complete(self, blocking: bool = True) -> Optional[WaitResult]: ... diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index 673d088e8443..7fb15a24acd4 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -53,6 +53,8 @@ SessionStatus, TxSessionBase, WaitResult, + derive_chunk_block_coords, + project_blocks_to_global_chunk, ) from tensorrt_llm._torch.disaggregation.native.auxiliary import ( AuxBuffer, @@ -98,6 +100,7 @@ class RecvReqInfo: dst_start_token: Optional[int] = None aux_slot: Optional[int] = None mamba_state_index: Optional[int] = None + # The receiver's own task index; the sender echoes it back in KV_AGENT_RESULT. slice_id: Optional[int] = None bounce_dst_base: Optional[int] = None @@ -152,7 +155,9 @@ class WriteMeta: dst_ptrs: np.ndarray # dtype=np.int64 sizes: np.ndarray # dtype=np.int64 dst_device_id: Optional[int] = None - slice_id: Optional[int] = None + sender_slice_id: Optional[int] = None + # The peer's task index, taken from RecvReqInfo.slice_id. + receiver_slice_id: int = 0 is_last_slice: bool = False meta_type: WriteMetaType = WriteMetaType.KV bounce_dst_base: Optional[int] = None @@ -181,15 +186,32 @@ class AgentResult(Enum): # KV_AGENT_RESULT prefix in one struct frame (was ascii frames serialized/parsed under the -# GIL per slice per writer): instance_rank, unique_rid, slice_id, is_last, status, -# transfer_size. The optional bounce tail follows at message[2:]. -_KV_RESULT_PREFIX = struct.Struct("gen wire format and has no version negotiation: the receiver unpacks +# whatever arrives against its compiled-in struct. Both servers must run matching builds. +_KV_RESULT_PREFIX = struct.Struct(" None: super().__init__(params) self.slice_id = slice_id self.transferred_count = 0 @@ -529,9 +562,10 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): logger.error(msg) write_meta.task.fail(RuntimeError(msg)) return - assert write_meta.slice_id is not None - task = session.kv_tasks[write_meta.slice_id] + assert write_meta.sender_slice_id is not None + task = session.kv_tasks[write_meta.sender_slice_id] timer = task._perf_timer + if timer: timer.record_push_end(write_meta.peer_rank) # Hold session.lock to serialize the INIT→TRANSFERRING transition with @@ -551,19 +585,12 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): f"in {status.value} state; sending FAILED to receiver" ) # Task may have been enqueued after cancel() already iterated kv_tasks, - # so its future was never set by cancel(). Set it here as a fallback. + # so its event was never set by cancel(). Set it here as a fallback. task.fail( RuntimeError(f"session {write_meta.unique_rid} {status.value}, transfer aborted") ) - self._get_or_connect_dealer(write_meta.peer_endpoint).send( - _make_kv_result_msg( - self._instance_rank, - write_meta.unique_rid, - write_meta.slice_id, - True, # is_last_slice — ensures receiver resolves its task future - AgentResult.FAILED, - ) - ) + # is_last=True ensures the receiver resolves its task event. + self._send_kv_result_to_receiver(write_meta, is_last=True, result=AgentResult.FAILED) return from .bounce import build_send_request, encode_result_tail @@ -582,17 +609,13 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): # region leak. Tell the receiver it failed and fail the local task instead. logger.error( f"_deliver_kv_to_agent: failed to build the KV send request for " - f"{write_meta.unique_rid} slice={write_meta.slice_id}: {e}" + f"{write_meta.unique_rid} sender_slice_id={write_meta.sender_slice_id}: {e}" ) task.fail(RuntimeError(f"build_send_request failed: {e}")) - self._get_or_connect_dealer(write_meta.peer_endpoint).send( - _make_kv_result_msg( - self._instance_rank, - write_meta.unique_rid, - write_meta.slice_id, - True, # is_last_slice — ensures receiver resolves its task future - AgentResult.FAILED, - ) + self._send_kv_result_to_receiver( + write_meta, + is_last=True, # ensures receiver resolves its task future + result=AgentResult.FAILED, ) return if timer: @@ -606,7 +629,7 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): detail = ( f"KV transfer agent failed: " f"unique_rid={write_meta.unique_rid} " - f"slice={write_meta.slice_id} " + f"sender_slice_id={write_meta.sender_slice_id} " f"peer_rank={write_meta.peer_rank} " f"peer_endpoint={write_meta.peer_endpoint} " f"op={getattr(request, 'op', '?')} " @@ -623,23 +646,21 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): if timer: timer.record_transfer_end(write_meta.peer_rank) - ## TODO: just last slice need to send task state? + # Intermediate chunk results are sent (not suppressed) so that RDMA + # failures propagate to the receiver immediately rather than requiring + # a timeout. Only the last chunk carries is_last_slice=True. tail = ( encode_result_tail(write_meta) if send_slot_id is not None and agent_result == AgentResult.SUCCESS else None ) - transfer_size = timer.get_transfer_size(write_meta.peer_rank) if timer else 0 - result_msg = _make_kv_result_msg( - self._instance_rank, - write_meta.unique_rid, - write_meta.slice_id, - write_meta.is_last_slice, - agent_result, - transfer_size=transfer_size, + self._send_kv_result_to_receiver( + write_meta, + is_last=write_meta.is_last_slice, + result=agent_result, + transfer_size=timer.get_transfer_size(write_meta.peer_rank) if timer else 0, tail=tail, ) - self._get_or_connect_thread_dealer(write_meta.peer_endpoint).send(result_msg) if timer: timer.record_task_end(write_meta.peer_rank) @@ -652,13 +673,14 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): if count > write_meta.expected_transfers: session.set_exception( - f"KV slice {write_meta.slice_id} received more than {write_meta.expected_transfers} transfers" + f"KV sender slice {write_meta.sender_slice_id} received more than " + f"{write_meta.expected_transfers} transfers" ) elif count == write_meta.expected_transfers: if task.is_done: task.status = TaskStatus.ERROR session.set_exception( - f"KV slice {write_meta.slice_id} task already resolved on completion" + f"KV sender slice {write_meta.sender_slice_id} task already resolved on completion" ) else: task.complete() @@ -667,9 +689,39 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): logger.debug( f"deliver_kv_to_agent completed: unique_rid={write_meta.unique_rid}, " - f"slice_id={write_meta.slice_id}, agent_result={agent_result}" + f"sender_slice_id={write_meta.sender_slice_id}, " + f"receiver_slice_id={write_meta.receiver_slice_id}, agent_result={agent_result}" ) + def _send_kv_result_to_receiver( + self, + write_meta: WriteMeta, + is_last: bool, + result: AgentResult, + transfer_size: int = 0, + tail=None, + ) -> None: + """Send a KV_AGENT_RESULT for a worker-thread delivery outcome. + + Covers every per-slice outcome (pre-transfer abort, RDMA failure, and + success). Sender-side chunking is transparent to the receiver because + the result addresses the receiver by its own task index + (``receiver_slice_id``, echoed from ``RecvReqInfo``); the sender's chunk + index rides along for logging and cross-side correlation. + Uses the per-thread DEALER because this runs on worker threads. + """ + result_msg = _make_kv_result_msg( + self._instance_rank, + write_meta.unique_rid, + write_meta.sender_slice_id if write_meta.sender_slice_id is not None else NO_SLICE_ID, + write_meta.receiver_slice_id, + is_last, + result, + transfer_size=transfer_size, + tail=tail, + ) + self._get_or_connect_thread_dealer(write_meta.peer_endpoint).send(result_msg) + @nvtx_range("_deliver_aux_to_agent") def _deliver_aux_to_agent(self, write_meta: WriteMeta): session = self._get_session(write_meta.unique_rid) @@ -826,14 +878,48 @@ def _build_kv_write_meta(self, task: KVSendTask, req_info: RecvReqInfo) -> Write # Aggregate fragments from all matching pools using numpy concatenation. # Send ownership is per pool: replicated pools elect one fan-in # owner, sharded pools keep head-duplication routing. + tpb = extractor.page_table.tokens_per_block + token_range = task._slice.token_range + slice_end = task._prompt_len + total_blocks = task._slice.total_blocks + if total_blocks is None: + total_blocks = (slice_end + tpb - 1) // tpb + # A slice is chunked (pipelined) when it is not the last slice or + # its block-aligned token_range starts past the sequence beginning. + # Only chunked slices need block-space chunk coordinates; a full + # transfer's token_range (end = prompt_len) need not be block-aligned, + # so avoid deriving/validating it in that case. + is_chunked = (not task._slice.is_last_slice) or ( + token_range is not None and token_range.start > 0 + ) + chunk_offset, chunk_block_count = ( + derive_chunk_block_coords(token_range, tpb) if is_chunked else (0, 0) + ) + # Resident block lists are the suffix of a range ending at the + # current global chunk boundary when pipelined, or at the full + # prompt end otherwise. token_start = (suffix_end - n_blocks) * tpb. + suffix_end_blocks = chunk_offset + chunk_block_count if is_chunked else total_blocks + for (self_lg, self_pi), (peer_lg, peer_pi) in pool_mapping.items(): if not self._registrar.should_send_pool(targets, peer_ri, self_lg, self_pi): continue src_block_ids = src_block_ids_per_groups[self_lg] - dst_block_ids = dst_block_ids_per_groups[peer_lg] + full_dst_block_ids = dst_block_ids_per_groups[peer_lg] + + # When sender uses chunking, the receiver sends all dst blocks + # in a single RecvReqInfo. Project the global chunk cursor into + # each destination layer group's resident/windowed block range. + if is_chunked: + dst_projectable_blocks = full_dst_block_ids[:total_blocks] + dst_block_ids = project_blocks_to_global_chunk( + dst_projectable_blocks, + chunk_block_offset=chunk_offset, + chunk_block_count=chunk_block_count, + resident_block_end=total_blocks, + ) + else: + dst_block_ids = full_dst_block_ids - tpb = extractor.page_table.tokens_per_block - token_range = task._slice.token_range lg_info = extractor.page_table.layer_groups[self_lg] window_size = getattr(lg_info, "sliding_window_size", None) @@ -845,10 +931,8 @@ def _build_kv_write_meta(self, task: KVSendTask, req_info: RecvReqInfo) -> Write beam_width=task._beam_width, ) - # Block lists are the suffix of [..., slice_end); cached prefix - # is implicit in their size. token_start = (total_blocks - n) * tpb. - slice_end = token_range.end if token_range is not None else 0 - total_blocks = (slice_end + tpb - 1) // tpb + # Cached prefix is implicit in the block-list size, so the token + # start is derived from suffix_end_blocks computed above. src_beam0_blocks = Sender._beam0_block_count( src_block_ids, total_blocks, task._beam_width ) @@ -863,19 +947,14 @@ def _build_kv_write_meta(self, task: KVSendTask, req_info: RecvReqInfo) -> Write f"dst beam-0 block list ({dst_beam0_blocks}) exceeds total slice " f"blocks ({total_blocks}); slice_end={slice_end}, tpb={tpb}" ) - src_start = (total_blocks - src_beam0_blocks) * tpb - dst_start = (total_blocks - dst_beam0_blocks) * tpb + src_start = (suffix_end_blocks - src_beam0_blocks) * tpb + dst_start = (suffix_end_blocks - dst_beam0_blocks) * tpb if req_info.dst_start_token is not None: dst_start = max(dst_start, req_info.dst_start_token) if window_size is not None: - # SWA stale_end uses the request prompt_len (not slice_end — - # they differ for non-final slices). prompt_len must be plumbed - # via the session; falling back to slice_end is wrong on - # non-final slices. - assert task._prompt_len is not None, ( - "SWA layer requires session.prompt_len; " - "set TxSession(prompt_len=request.prompt_len)." - ) + # The stale-block boundary is a property of the whole request, + # so it uses prompt_len rather than this chunk's end; the two + # differ on every non-final slice. stale_end = max(0, (task._prompt_len + 1 - window_size) // tpb) src_start = max(stale_end * tpb, src_start) dst_start = max(stale_end * tpb, dst_start) @@ -942,7 +1021,8 @@ def _build_kv_write_meta(self, task: KVSendTask, req_info: RecvReqInfo) -> Write peer_rank=peer_ri.instance_rank, peer_endpoint=peer_ri.self_endpoint, unique_rid=task._unique_rid, - slice_id=task.slice_id, + sender_slice_id=task.slice_id, + receiver_slice_id=req_info.slice_id if req_info.slice_id is not None else 0, is_last_slice=task._slice.is_last_slice, bounce_dst_base=req_info.bounce_dst_base, ) @@ -1082,39 +1162,41 @@ def _handle_cancel_session(self, message: list[bytes]): @nvtx_range("_respond_with_kv") def _respond_with_kv(self, _send_id: bytes, message: list[bytes]): # _sessions_lock prevents a race between session lookup and req_info save. - # session.lock serializes _enqueue calls from both paths. + # dispatch_lock keeps replayed tasks ahead of concurrently added slices. info: RecvReqInfo = RecvReqInfo.from_bytes(message[1]) with self._sessions_lock: session = self._get_session(info.unique_rid) if session is None: self._save_peer_req_info(info) return - with session.lock: - self._save_peer_req_info(info) - tasks = list(session.kv_tasks) - # No tasks: no worker will send KV_AGENT_RESULT FAILED to the receiver. - # Send it directly to unblock the receiver's TRANSFERRING task future; - # CANCEL_SESSION alone would leave it stuck indefinitely. - if not tasks and session.status in (SessionStatus.ERROR, SessionStatus.CANCELLED): - self._send_failed_result_to_receiver(info) - return - for task in tasks: - if task._perf_timer is not None: - task._perf_timer.record_task_start(info.instance_rank) - trans_meta = self._build_kv_write_meta(task, info) - if task._perf_timer is not None: - task._perf_timer.record_push_start(trans_meta.peer_rank) - self._enqueue(trans_meta) + with session.dispatch_lock: + with session.lock: + self._save_peer_req_info(info) + tasks = list(session.kv_tasks) + # No tasks: no worker will send KV_AGENT_RESULT FAILED to the receiver. + # Send it directly to unblock the receiver's TRANSFERRING task event; + # CANCEL_SESSION alone would leave it stuck indefinitely. + if not tasks and session.status in (SessionStatus.ERROR, SessionStatus.CANCELLED): + self._send_failed_result_to_receiver(info) + return + for task in tasks: + if task._perf_timer is not None: + task._perf_timer.record_task_start(info.instance_rank) + trans_meta = self._build_kv_write_meta(task, info) + if task._perf_timer is not None: + task._perf_timer.record_push_start(trans_meta.peer_rank) + self._enqueue(trans_meta) def _send_failed_result_to_receiver(self, info: RecvReqInfo): try: peer_ri = self._registrar.get_peer_rank_info(info.instance_name, info.instance_rank) - slice_id = info.slice_id if info.slice_id is not None else 0 + receiver_slice_id = info.slice_id if info.slice_id is not None else 0 self._get_or_connect_dealer(peer_ri.self_endpoint).send( _make_kv_result_msg( self._instance_rank, info.unique_rid, - slice_id, + NO_SLICE_ID, # no KVSendTask owns this result + receiver_slice_id, True, # is_last_slice AgentResult.FAILED, ) @@ -1227,9 +1309,9 @@ def __init__( request_id: int, params: DisaggregatedParams, sender: Sender, + prompt_len: int, aux_buffer: Optional[AuxBuffer] = None, timeout_s: Optional[float] = None, - prompt_len: Optional[int] = None, beam_width: int = 1, ): super().__init__( @@ -1246,6 +1328,10 @@ def __init__( self.kv_tasks = [] self.aux_task = None self.lock = threading.Lock() + # Serializes task discovery and queue insertion between live sends and + # late-peer replay. The worker FIFO makes is_last_slice correct only + # when every older slice is enqueued first. + self.dispatch_lock = threading.Lock() self._exception: Optional[Exception] = None self._closed = False @@ -1272,6 +1358,10 @@ def disagg_request_id(self) -> int: def status(self) -> SessionStatus: if self._terminal_status is not None: return self._terminal_status + if self._exception is not None or any(t.status == TaskStatus.ERROR for t in self.kv_tasks): + return SessionStatus.ERROR + if self.aux_task is not None and self.aux_task.status == TaskStatus.ERROR: + return SessionStatus.ERROR kv_all_transferred = bool(self.kv_tasks) and all( t.status == TaskStatus.TRANSFERRED for t in self.kv_tasks ) @@ -1286,29 +1376,31 @@ def status(self) -> SessionStatus: def send(self, slice: KVSlice) -> None: if self.transfer_start_time is None: self.transfer_start_time = tensorrt_llm.bindings.global_steady_clock_now() - with self.lock: - params = self._base_args.params - slice_id = len(self.kv_tasks) - task = KVSendTask( - slice, - params, - slice_id, - prompt_len=self._base_args.prompt_len, - beam_width=self._base_args.beam_width, - ) - task._unique_rid = self.disagg_request_id - self.kv_tasks.append(task) - req_info_snapshot = dict(self._sender._get_req_info(task._unique_rid) or {}) - self._sender.dispatch_task(task, req_info_snapshot) + with self.dispatch_lock: + with self.lock: + params = self._base_args.params + slice_id = len(self.kv_tasks) + task = KVSendTask( + slice, + params, + slice_id, + prompt_len=self._base_args.prompt_len, + beam_width=self._base_args.beam_width, + ) + task._unique_rid = self.disagg_request_id + self.kv_tasks.append(task) + req_info_snapshot = dict(self._sender._get_req_info(task._unique_rid) or {}) + self._sender.dispatch_task(task, req_info_snapshot) def send_aux(self) -> AuxSendTask: - with self.lock: - params = self._base_args.params - task = AuxSendTask(params, self.aux_slot) - task._unique_rid = self.disagg_request_id - self.aux_task = task - req_info_snapshot = dict(self._sender._get_req_info(task._unique_rid) or {}) - self._sender.dispatch_task(task, req_info_snapshot) + with self.dispatch_lock: + with self.lock: + params = self._base_args.params + task = AuxSendTask(params, self.aux_slot) + task._unique_rid = self.disagg_request_id + self.aux_task = task + req_info_snapshot = dict(self._sender._get_req_info(task._unique_rid) or {}) + self._sender.dispatch_task(task, req_info_snapshot) return task def pack_aux(self, request: LlmRequest) -> None: @@ -1815,9 +1907,15 @@ def _process_kv_agent_result(self, _send_id: bytes, message: list[bytes]): f"_process_kv_agent_result: unexpected msg_type={message[0]!r}, expected KV_AGENT_RESULT" ) return - peer_rank, unique_rid, sender_slice_id, is_last_slice, status_code, transfer_size = ( - _KV_RESULT_PREFIX.unpack(message[1]) - ) + ( + peer_rank, + unique_rid, + sender_slice_id, + receiver_slice_id, + is_last_slice, + status_code, + transfer_size, + ) = _KV_RESULT_PREFIX.unpack(message[1]) from .bounce import decode_result_tail dst_ptrs, sizes, src_base = decode_result_tail(message) @@ -1829,6 +1927,7 @@ def _process_kv_agent_result(self, _send_id: bytes, message: list[bytes]): return session.process_kv_agent_result( peer_rank, + receiver_slice_id, sender_slice_id, is_last_slice, _AGENT_RESULT_BY_CODE[status_code], @@ -1875,9 +1974,9 @@ def __init__( request_id: int, params: DisaggregatedParams, receiver: Receiver, + prompt_len: int, aux_buffer: Optional[AuxBuffer] = None, timeout_s: Optional[float] = None, - prompt_len: Optional[int] = None, beam_width: int = 1, ): super().__init__( @@ -1953,6 +2052,7 @@ def receive(self, slice: KVSlice) -> None: def process_kv_agent_result( self, peer_rank: int, + receiver_slice_id: int, sender_slice_id: int, is_last_slice: bool, status: AgentResult, @@ -1961,14 +2061,21 @@ def process_kv_agent_result( src_base=None, transfer_size: int = 0, ): + """Apply one sender chunk's delivery outcome to this session. + + Args: + receiver_slice_id: Index of the local task this result resolves. + sender_slice_id: The sender's chunk index, for logging only; + NO_SLICE_ID when the sender had no task to attribute it to. + """ with self.lock: self.kv_cache_size_bytes += transfer_size - assert sender_slice_id < len(self._kv_tasks), ( - f"Receiver got slice_id={sender_slice_id} from sender but only has " - f"{len(self._kv_tasks)} receive task(s) for request {self.request_id}. " - f"Sender/receiver slice count mismatch." + assert receiver_slice_id < len(self._kv_tasks), ( + f"Receiver got receiver_slice_id={receiver_slice_id} (sender_slice_id=" + f"{sender_slice_id}) but only has {len(self._kv_tasks)} receive task(s) " + f"for request {self.request_id}. Sender/receiver slice count mismatch." ) - task = self._kv_tasks[sender_slice_id] + task = self._kv_tasks[receiver_slice_id] if status == AgentResult.SUCCESS: from .bounce import scatter_write_result @@ -1988,6 +2095,7 @@ def on_done( success, task=task, peer_rank=peer_rank, + receiver_slice_id=receiver_slice_id, sender_slice_id=sender_slice_id, request_id=request_id, instance_name=instance_name, @@ -2001,7 +2109,8 @@ def on_done( task.fail( RuntimeError( f"KV bounce scatter failed for request {request_id} " - f"slice={sender_slice_id}" + f"receiver_slice_id={receiver_slice_id} " + f"sender_slice_id={sender_slice_id}" ) ) return @@ -2014,7 +2123,8 @@ def on_done( except Exception as e: # perf is best-effort; never block completion logger.warning( f"KV transfer perf logging failed for request {request_id} " - f"slice={sender_slice_id}: {e}" + f"receiver_slice_id={receiver_slice_id} " + f"sender_slice_id={sender_slice_id}: {e}" ) task.complete() # Transfer end for perf/time-sync: only meaningful once every slice has @@ -2026,7 +2136,8 @@ def on_done( ) logger.debug( f"KV transfer complete for request {request_id} " - f"slice={sender_slice_id}" + f"receiver_slice_id={receiver_slice_id} " + f"sender_slice_id={sender_slice_id}" ) scatter_write_result( @@ -2040,7 +2151,8 @@ def on_done( ) elif status == AgentResult.FAILED: detail = ( - f"KV transfer failed for request {self.request_id} slice={sender_slice_id} " + f"KV transfer failed for request {self.request_id} " + f"receiver_slice_id={receiver_slice_id} sender_slice_id={sender_slice_id} " f"peer_rank={peer_rank} is_last_slice={is_last_slice} " f"(reported by remote agent; see sender-side log for nixl_status)" ) @@ -2059,8 +2171,8 @@ def on_done( ) def process_aux_agent_result(self, _peer_rank: int, status: AgentResult): - # Aux is session-level (not per-slice); expected_transfers is identical - # across all kv_tasks, so any task provides the right count. + # Aux is session-level (not per-slice); the RxSession's single task's + # expected_transfers is the number of sender transfers to wait for. with self.lock: if not self._kv_tasks: logger.warning( @@ -2336,7 +2448,18 @@ def populate_instance_and_rank_info(self, endpoints: list[str], layer_num_per_pp self._rank_info.sender_endpoints = endpoints self._rank_info.layer_num_per_pp = layer_num_per_pp - def create_tx_session(self, request: LlmRequest) -> TxSession: + def create_tx_session( + self, + request: LlmRequest, + ) -> TxSession: + """Create a TxSession for the given request. + + Args: + request: The LLM request to create a send session for. + + Returns: + A new ``TxSession`` ready to accept ``send()`` calls. + """ params = request.py_disaggregated_params assert params is not None return TxSession( diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index 0d7e3429424b..9bae66d90506 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -32,6 +32,7 @@ TxSessionBase, WaitResult, get_unique_rid, + project_blocks_to_global_chunk, ) from tensorrt_llm._torch.disaggregation.native.bounce import ( config_from_size as bounce_config_from_size, @@ -143,6 +144,7 @@ def __init__( # _slice_num_bytes() is this rank's KV shard, so scale by tp_size to get the request total (kv_cache_size), # except under attention DP where the local count already is the total. self._kv_size_rank_factor = 1 if mapping.enable_attention_dp else max(1, mapping.tp_size) + self._enable_pipelined_transfer = cache_transceiver_config.enable_pipelined_transfer # Sticky role markers; flip True once any session opens, used to short-circuit # per-iter tp_allgather when this transceiver never sends/receives. @@ -250,11 +252,30 @@ def __enter__(self): def __exit__(self, _exc_type, _exc_val, _exc_tb): self.shutdown() - def _create_kv_slice(self, req: LlmRequest) -> KVSlice: + def _create_kv_slice( + self, + req: LlmRequest, + resident_block_end: Optional[int] = None, + ) -> KVSlice: + """Create a KV slice from the source's currently resident blocks. + + Args: + req: Request whose KV blocks are being described. + resident_block_end: Exclusive logical block boundary to include. + Pipelined prefill passes the current chunk end to exclude + full-prompt blocks that V1 reserved but has not computed yet. + ``None`` includes the complete prompt. + """ adapter = self._reuse_adapter tpb = adapter.tokens_per_block assert self._page_table is not None layer_groups = self._page_table.layer_groups + prompt_blocks = (req.prompt_len + tpb - 1) // tpb + resident_blocks = ( + prompt_blocks + if resident_block_end is None + else min(max(0, resident_block_end), prompt_blocks) + ) is_gen_only = req.is_generation_only_request() cached_per_lg = ( @@ -283,29 +304,27 @@ def _create_kv_slice(self, req: LlmRequest) -> KVSlice: continue block_ids = adapter.get_block_ids(req, idx, lg) # Limit to prompt_len blocks, matching C++ cacheFormatter behavior. - total_blocks = (req.prompt_len + tpb - 1) // tpb - if block_ids.size > total_blocks: - block_ids = block_ids[:total_blocks] + if block_ids.size > resident_blocks: + block_ids = block_ids[:resident_blocks] window_size = lg.sliding_window_size if window_size is not None: # Drop stale blocks the manager may still expose (V1 pre-eviction). stale_end = max(0, (req.prompt_len + 1 - window_size) // tpb) - expected_valid = max(0, total_blocks - stale_end) + expected_valid = max(0, resident_blocks - stale_end) # Stale prefix already pruned above; skip reuse-hit blocks that # land inside the window. Clamp to 0: ctx side has cached_per_lg # synthetically 0, and a reuse hit may fall entirely inside the # stale region (those blocks were already pruned, no extra skip). cache_skip = max(0, cached_per_lg[idx] // tpb - stale_end) else: - total_blocks = (req.prompt_len + tpb - 1) // tpb - expected_valid = total_blocks + expected_valid = resident_blocks cache_skip = cached_per_lg[idx] // tpb block_ids = self._trim_packed_beam_block_ids( block_ids, beam_width=req.py_beam_width, - total_blocks=total_blocks, + total_blocks=resident_blocks, expected_valid=expected_valid, cache_skip=cache_skip, ) @@ -579,13 +598,34 @@ def _build_to_process( to_process.append(rid) return to_process - def _close_failed_sessions(self, sessions: dict, reqs: dict, failed: list): + def _close_failed_sessions( + self, sessions: dict, reqs: dict, failed: list, mark_retired: bool = False + ): for rid in failed: reqs[rid].state = LlmRequestState.DISAGG_TRANS_ERROR + if mark_retired: + reqs[rid].py_kv_send_session_retired = True sessions[rid].close() del reqs[rid] del sessions[rid] + def _retire_send_session(self, rid: int, req: Optional[LlmRequest] = None) -> None: + """Close a send session and bar any later chunk from re-creating it. + + close() drops the peer's RecvReqInfo (Sender.clear_session) and the + receiver never re-sends it, so a session created after this point has + no peer to write to and would leave every task in INIT. + """ + session = self._send_sessions.pop(rid, None) + if session is not None: + session.close() + # _send_reqs is only populated once a slice has been built, so an early + # teardown has to be handed the request explicitly. + req = req if req is not None else self._send_reqs.get(rid) + self._send_reqs.pop(rid, None) + if req is not None: + req.py_kv_send_session_retired = True + def _apply_aux(self, session, req: LlmRequest): """Unpack aux tokens from session into request's context_phase_params.""" session.unpack_aux(req) @@ -605,10 +645,18 @@ def _apply_aux(self, session, req: LlmRequest): req.context_phase_params.first_gen_tokens = first_gen_tokens req.context_phase_params.draft_tokens = draft_tokens - def _get_or_create_send_session(self, req: LlmRequest) -> TxSessionBase: + def _get_or_create_send_session(self, req: LlmRequest) -> Optional[TxSessionBase]: + self._ever_had_send_session = True rid = get_unique_rid(req) assert rid is not None if rid not in self._send_sessions: + if req.py_kv_send_session_retired: + logger.warning( + f"rid={rid}: send session already retired; failing the request " + "rather than re-creating one with no peer registration" + ) + req.state = LlmRequestState.DISAGG_TRANS_ERROR + return None self._send_sessions[rid] = self._transfer_worker.create_tx_session(req) return self._send_sessions[rid] @@ -629,14 +677,125 @@ def _finalize_send(self, req: LlmRequest, session: TxSessionBase): ) self._send_reqs[rid] = req + @property + def pipeline_transfer_enabled(self) -> bool: + """Whether pipelined prefill-transfer is enabled.""" + return self._enable_pipelined_transfer + + def has_inflight_transfer(self, req: LlmRequest) -> bool: + """Whether transfer resources are still owned for this request. + + Session membership is the ownership record: a session is registered + before its first slice is sent and removed only once the transfer is + complete, cancelled, or failed. This is deliberately independent of + ``req.state``, which tracks the compute phase and lags behind the + transfer during pipelined prefill. + """ + rid = get_unique_rid(req) + return rid in self._send_sessions or rid in self._recv_sessions + + def has_any_inflight_transfer(self) -> bool: + """Whether any request has transfer resources in flight.""" + return bool(self._send_sessions) or bool(self._recv_sessions) + + def has_retired_send_session(self, req: LlmRequest) -> bool: + """Whether req's send session was torn down before its last slice.""" + return req.py_kv_send_session_retired and get_unique_rid(req) not in self._send_sessions + + def _build_prefill_chunk( + self, + req: LlmRequest, + ) -> KVSlice: + """ + Create a KVSlice for a prefill chunk. Project the block IDs to the global chunk. + + Args: + req: The context-only request being prefilled. + """ + assert req.py_beam_width == 1, "beam_width > 1 is not supported for chunked KV transfer" + rid = get_unique_rid(req) + assert rid is not None + self._send_reqs[rid] = req + + chunk_start_pos, chunk_end_pos = req.py_last_context_chunk + tpb = self._kv_cache_manager.tokens_per_block + + # A ctx-side prefix-reuse hit starts the first chunk at + # prepopulated_prompt_len, so no chunk covers [0, prepopulated_prompt_len). + # Those blocks are resident and valid and the generation server still needs + # them, so the first slice extends back to block 0. req.is_first_context_chunk + # cannot be used here: it compares context_current_position against + # prepopulated_prompt_len, and _update_request_states has already advanced the + # cursor by the time _send_kv_async runs. The recorded chunk start is the + # pre-advance value. + is_first_chunk = chunk_start_pos == req.prepopulated_prompt_len + chunk_start_block = 0 if is_first_chunk else chunk_start_pos // tpb + chunk_end_block = (chunk_end_pos + tpb - 1) // tpb + is_last_chunk = req.context_remaining_length == 0 + + prompt_blocks = (req.prompt_len + tpb - 1) // tpb + total_blocks = prompt_blocks + + chunk_start = min(chunk_start_block, total_blocks) + chunk_end = min(chunk_end_block, total_blocks) + chunk_block_count = max(0, chunk_end - chunk_start) + # V1 reserves the full prompt up front, while V2 grows its source list + # incrementally. Normalize both to the current chunk boundary before + # projecting so full-prompt SWA pages are never treated as current pages. + base_slice = self._create_kv_slice(req, resident_block_end=chunk_end) + all_block_ids = base_slice.block_ids_per_layer_groups + chunk_block_ids = [ + project_blocks_to_global_chunk( + block_ids, + chunk_block_offset=chunk_start, + chunk_block_count=chunk_block_count, + resident_block_end=chunk_end, + ) + for block_ids in all_block_ids + ] + chunk_token_range = None + if chunk_block_count > 0: + chunk_token_range = TokenRange( + start=chunk_start * tpb, + end=(chunk_start + chunk_block_count) * tpb, + ) + + return KVSlice( + is_last_slice=is_last_chunk, + block_ids_per_layer_groups=chunk_block_ids, + mamba_state_index=base_slice.mamba_state_index, + token_range=chunk_token_range, + total_blocks=total_blocks, + ) + @nvtx_range("KvCacheTransceiverV2.respond_and_send_async") - def respond_and_send_async(self, req: LlmRequest): - self._ever_had_send_session = True + def respond_and_send_async(self, req: LlmRequest) -> None: + """Start background KV cache transfer to the generation server. + + Creates (or reuses) a ``TxSession`` and sends a KV slice (monolithic or chunked) for each request. + + Args: + req: The completed context request whose KV cache to transfer. + """ + + # Pipelined transfer records the transfer start time of the last slice. + # The records of the previous slices are overwritten. req.set_kv_cache_transfer_start(tensorrt_llm.bindings.global_steady_clock_now()) session = self._get_or_create_send_session(req) - req.state = LlmRequestState.DISAGG_CONTEXT_TRANS_IN_PROGRESS - session.send(self._create_kv_slice(req)) - self._finalize_send(req, session) + if session is None: + return + if self.pipeline_transfer_enabled: + slice = self._build_prefill_chunk(req) + else: + slice = self._create_kv_slice(req) + session.send(slice) + + if slice.is_last_slice: + self._finalize_send(req, session) + # This marks the compute-phase boundary only. Transfer ownership + # began at the first session.send() above and is tracked by session + # membership (has_inflight_transfer), not by req.state. + req.state = LlmRequestState.DISAGG_CONTEXT_TRANS_IN_PROGRESS @nvtx_range("KvCacheTransceiverV2.request_and_receive_sync") def request_and_receive_sync(self, req: LlmRequest): @@ -675,7 +834,17 @@ def request_and_receive_sync(self, req: LlmRequest): self._recv_reqs.pop(rid, None) @nvtx_range("KvCacheTransceiverV2.request_and_receive_async") - def request_and_receive_async(self, req: LlmRequest): + def request_and_receive_async(self, req: LlmRequest) -> None: + """Start background KV cache receive from the context server. + + The receiver always uses a single monolithic slice. Chunking is + sender-only: the sender splits its source blocks into chunks and + slices the receiver's destination blocks to match each chunk. + + Args: + req: The generation request whose KV cache blocks to receive + into. + """ self._ever_had_recv_session = True req.set_kv_cache_transfer_start(tensorrt_llm.bindings.global_steady_clock_now()) rid = get_unique_rid(req) @@ -703,6 +872,7 @@ def check_context_transfer_status( if self._ctx_need_tp_sync or self._ctx_need_pp_sync: self._transfer_worker.sweep_stale_req_infos() return [], [] + block_all = at_least_request_num is None wait_num = at_least_request_num if not block_all else 0 need_progress = wait_num > 0 @@ -747,17 +917,13 @@ def check_context_transfer_status( ) for rid in cancelled: - self._send_sessions[rid].close() - del self._send_reqs[rid] - del self._send_sessions[rid] + self._retire_send_session(rid) for rid in completed: if mark_complete: self._send_reqs[rid].state = LlmRequestState.DISAGG_CONTEXT_COMPLETE - self._send_sessions[rid].close() - del self._send_reqs[rid] - del self._send_sessions[rid] - self._close_failed_sessions(self._send_sessions, self._send_reqs, failed) + self._retire_send_session(rid) + self._close_failed_sessions(self._send_sessions, self._send_reqs, failed, mark_retired=True) # Sweep orphaned RecvReqInfo entries from ADP broadcast on non-assigned # DP ranks (entries that will never have a TxSession created for them). @@ -933,9 +1099,7 @@ def cancel_request(self, req: LlmRequest) -> bool: if self._send_sessions[rid].has_transferring_tasks(): has_transferring = True else: - self._send_sessions[rid].close() - del self._send_reqs[rid] - del self._send_sessions[rid] + self._retire_send_session(rid, req) if rid in self._recv_sessions: self._recv_sessions[rid].cancel() diff --git a/tensorrt_llm/_torch/pyexecutor/_util.py b/tensorrt_llm/_torch/pyexecutor/_util.py index c4fb115111fe..b2725a9caa5f 100644 --- a/tensorrt_llm/_torch/pyexecutor/_util.py +++ b/tensorrt_llm/_torch/pyexecutor/_util.py @@ -3006,8 +3006,13 @@ def create_py_executor_instance( mamba_cache_manager = kv_cache_manager kv_cache_transceiver = create_kv_cache_transceiver( - mapping, dist, kv_cache_manager, attention_type, - cache_transceiver_config, mamba_cache_manager) + mapping, + dist, + kv_cache_manager, + attention_type, + cache_transceiver_config, + mamba_cache_manager, + enable_chunked_prefill=llm_args.enable_chunked_prefill) waiting_queue_policy = (scheduler_config.waiting_queue_policy if scheduler_config is not None else diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py index d4868acf0226..e0419778c470 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py @@ -16,6 +16,7 @@ from .llm_request import LlmRequest from .mamba_cache_manager import (BaseMambaCacheManager, CppMambaHybridCacheManager, + MambaHybridCacheManager, MambaHybridCacheManagerV2, MixedMambaHybridCacheManager) from .resource_manager import KVCacheManager @@ -100,6 +101,64 @@ def _validate_disagg_inflight_cancel_config( f"layerwise={layerwise!r}, try_zcopy={try_zcopy!r}.") +def resolve_cache_transceiver_config( + cache_transceiver_config: Optional[CacheTransceiverConfig]) -> None: + """Resolve defaults and validate runtime-independent configuration.""" + if cache_transceiver_config is None or cache_transceiver_config.backend is None: + return + + # "auto" is normally resolved against the model's preference at config load + # time (ModelLoader.load_config_and_apply_defaults); paths that skip that + # step (e.g. AutoDeploy) fall back to the C++ transceiver. Collapse it to + # None here so every consumer below - including the pipelined-transfer + # auto-selection - sees a resolved runtime. + if cache_transceiver_config.transceiver_runtime == "auto": + cache_transceiver_config.transceiver_runtime = None + + # Resolved for the checks below only; "DEFAULT" is deliberately left on the + # config so that _validate_disagg_inflight_cancel_config() can still reject + # it as an ambiguous backend. create_kv_cache_transceiver() commits the + # resolved value after that validation runs. + effective_backend = cache_transceiver_config.backend + if effective_backend == "DEFAULT": + effective_backend, _ = cache_transceiver_config._resolve_default_backend( + ) + + runtime = cache_transceiver_config.transceiver_runtime + enable_pipelined_transfer = cache_transceiver_config.enable_pipelined_transfer + if runtime is None and enable_pipelined_transfer: + if effective_backend != "NIXL": + raise ValueError( + f"enable_pipelined_transfer is set but backend " + f"'{effective_backend}' requires the C++ " + f"transceiver, which does not support pipelined transfer. Use NIXL backend to " + f"enable pipelined transfer.") + logger.warning( + "enable_pipelined_transfer is set; auto-selecting the Python " + "transceiver instead of the C++ transceiver to enable " + "pipelined KV cache transfer. " + "Set transceiver_runtime='CPP' to disable this auto-selection.") + cache_transceiver_config.transceiver_runtime = "PYTHON" + elif runtime == "CPP" and enable_pipelined_transfer: + raise ValueError( + "enable_pipelined_transfer is set but transceiver_runtime='CPP' " + "explicitly disables Python auto-selection. Use transceiver_runtime='PYTHON' to enable pipelined transfer." + ) + + if (cache_transceiver_config.transceiver_runtime == "PYTHON" + and effective_backend != "NIXL"): + raise ValueError( + f"Python transceiver currently only supports NIXL backend, " + f"got {effective_backend}. " + f"Please use transceiver_runtime='CPP' for MPI, UCX, or MOONCAKE backends." + ) + if (enable_pipelined_transfer + and cache_transceiver_config.kv_cache_bounce_size_mb > 0): + raise ValueError( + "kv_cache_bounce_size_mb must be 0 when enable_pipelined_transfer is set." + ) + + def mapping_to_world_config(mapping: Mapping) -> WorldConfig: return WorldConfig(tensor_parallelism=mapping.tp_size, @@ -117,26 +176,22 @@ def create_kv_cache_transceiver( kv_cache_manager: KVCacheManager, attention_type: AttentionTypeCpp, cache_transceiver_config: CacheTransceiverConfig, - mamba_cache_manager: Optional[BaseMambaCacheManager] = None): + mamba_cache_manager: Optional[BaseMambaCacheManager] = None, + enable_chunked_prefill: bool = False): + resolve_cache_transceiver_config(cache_transceiver_config) if cache_transceiver_config is None or cache_transceiver_config.backend is None: logger.info("cache_transceiver is disabled") return None - # "auto" is normally resolved against the model's preference at config - # load time (ModelLoader.load_config_and_apply_defaults); paths that skip - # that step (e.g. AutoDeploy) fall back to the C++ transceiver here. This - # must run before any consumer of transceiver_runtime below (e.g. the - # inflight-cancel validation, which treats non-CPP runtimes as - # unsupported). - if cache_transceiver_config.transceiver_runtime == "auto": - cache_transceiver_config.transceiver_runtime = None - + # transceiver_runtime is already resolved by resolve_cache_transceiver_config + # above; these checks need the cache managers, so they cannot move there. if (cache_transceiver_config.transceiver_runtime != "PYTHON" and isinstance(mamba_cache_manager, MixedMambaHybridCacheManager)): raise ValueError( "MixedMambaHybridCacheManager requires the Python transceiver " "runtime in disaggregated serving.") + # Runs while backend may still be "DEFAULT", which it rejects as ambiguous. _validate_disagg_inflight_cancel_config(cache_transceiver_config) if cache_transceiver_config.backend == "DEFAULT": @@ -158,6 +213,24 @@ def create_kv_cache_transceiver( "UCX_CUDA_IPC_ENABLE_MNNVL=n, UCX_RNDV_SCHEME=put_zcopy and/or unset UCX_NET_DEVICES upon server " "hangs or lower-than-expected performance.") + if (cache_transceiver_config.enable_pipelined_transfer + and (isinstance(kv_cache_manager, MambaHybridCacheManager) + or mamba_cache_manager is not None)): + raise ValueError( + "enable_pipelined_transfer is not supported with Mamba/hybrid attention models." + ) + if (cache_transceiver_config.enable_pipelined_transfer + and not enable_chunked_prefill): + raise ValueError( + "enable_chunked_prefill is required when enable_pipelined_transfer is set." + ) + is_kv_cache_sender = getenv("TRTLLM_DISAGG_ROLE") != "generation" + if (cache_transceiver_config.enable_pipelined_transfer + and is_kv_cache_sender and mapping.pp_size != 1): + raise ValueError( + "pipeline_parallel_size=1 is required when enable_pipelined_transfer is set." + ) + # Select transceiver implementation based on transceiver_runtime. # transceiver_runtime == None or "CPP" -> use C++ transceiver (default) # transceiver_runtime == "PYTHON" -> use Python transceiver. @@ -191,6 +264,7 @@ def create_kv_cache_transceiver( from tensorrt_llm._torch.disaggregation.transceiver import \ KvCacheTransceiverV2 logger.info("Using KvCacheTransceiverV2") + # MixedMambaHybridCacheManager contains both the KV and Mamba pools. return KvCacheTransceiverV2(mapping, dist, kv_cache_manager, cache_transceiver_config) @@ -202,6 +276,34 @@ def create_kv_cache_transceiver( class KvCacheTransceiver(ABC): + @property + def pipeline_transfer_enabled(self) -> bool: + """Whether pipelined prefill-transfer is enabled.""" + return False + + def has_inflight_transfer(self, req: LlmRequest) -> bool: + """Whether this transceiver still owns transfer resources for req. + + Independent of ``LlmRequestState``: with pipelined transfer a chunk can + be in flight while the request is still in its context-compute phase. + True means the request's KV pages may be read by the fabric and must + not be released. + """ + return False + + def has_any_inflight_transfer(self) -> bool: + """Whether any request has transfer resources in flight.""" + return False + + def has_retired_send_session(self, req: LlmRequest) -> bool: + """Whether req's send session was torn down before its last slice. + + Tearing it down also drops the peer's receive registration, so no + further slice can reach the generation server. The request has to be + failed rather than fed another chunk. + """ + return False + @abstractmethod def respond_and_send_async(self, req: LlmRequest): raise NotImplementedError diff --git a/tensorrt_llm/_torch/pyexecutor/llm_request.py b/tensorrt_llm/_torch/pyexecutor/llm_request.py index c68b83c2abde..f8709a99f0cb 100644 --- a/tensorrt_llm/_torch/pyexecutor/llm_request.py +++ b/tensorrt_llm/_torch/pyexecutor/llm_request.py @@ -786,6 +786,10 @@ def __init__( self.is_cuda_graph_dummy = False self.py_kv_transfer_start_time = None self.py_kv_transfer_timed_out = False + # Set when the send session is torn down. Closing it also drops the + # peer's receive registration, which never comes back, so a session + # created after this point would have nobody to write to. + self.py_kv_send_session_retired = False # Encoder-decoder runtime state. ``py_encoder_output`` holds the # packed encoder hidden states produced by the encoder iteration as diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor.py b/tensorrt_llm/_torch/pyexecutor/py_executor.py index 62b7c6488ee9..bf7a661b3534 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor.py @@ -2539,7 +2539,7 @@ def _pp_retry_until_can_schedule(self, scheduled_batch): raise RuntimeError( "KV cache transceiver is not enabled, but current rank cannot run first PP's schedule result due to limited KV cache resources. This is not expected." ) - if not self.async_transfer_manager.has_any_inflight_requests(): + if not self._has_any_inflight_kv_transfer(): raise RuntimeError( "No context cache transmission is in progress, but current rank cannot run first PP's schedule result due to limited KV cache resources. This is not expected." ) @@ -3213,8 +3213,7 @@ def _handle_executed_batch(self, executed_batch: Optional[BatchStatePP]): self._remove_inflight_ids(scheduled_requests) - if self.kv_cache_transceiver and self.async_transfer_manager.has_any_inflight_requests( - ): + if self.kv_cache_transceiver and self._has_any_inflight_kv_transfer(): self._check_kv_transfer_timeout() if self._disagg_pp_termination_handler is not None: @@ -4251,8 +4250,8 @@ def _executor_loop(self): # collective fires every iter regardless of future restructuring. self._handle_kv_transfer_timeouts_synced() - if self.kv_cache_transceiver and self.async_transfer_manager.has_any_inflight_requests( - ): + if (self.kv_cache_transceiver + and self._has_any_inflight_kv_transfer()): self._check_kv_transfer_timeout() self._kv_connector_terminate_requests() @@ -4801,8 +4800,8 @@ def _executor_loop_overlap(self): # If the batch is empty on this rank, we need to clear the previous batch. self.previous_batch = None - if self.kv_cache_transceiver and self.async_transfer_manager.has_any_inflight_requests( - ): + if (self.kv_cache_transceiver + and self._has_any_inflight_kv_transfer()): self._check_kv_transfer_timeout() self._kv_connector_terminate_requests() @@ -4983,6 +4982,19 @@ def _validate_request(self, request: LlmRequest): # Check token ID ranges self._validate_token_id_range(request) + if (not self.is_warmup and self.kv_cache_transceiver is not None + and self.kv_cache_transceiver.pipeline_transfer_enabled): + if request.py_beam_width != 1: + raise ValueError( + "beam_width > 1 is not supported when enable_pipelined_transfer is set." + ) + + disagg_params = request.py_disaggregated_params + if (disagg_params is not None and disagg_params.schedule_style + != DisaggScheduleStyle.GENERATION_FIRST): + raise ValueError("schedule_style must be generation_first when " + "enable_pipelined_transfer is set.") + # Perform sampler-specific validation self.sampler.validate_request(request) @@ -5989,6 +6001,18 @@ def _check_gen_cache_transfer_errors_consensus(self) -> None: requests=error_requests, charge_budget=False) + def _has_any_inflight_kv_transfer(self) -> bool: + """Whether any KV transfer resources are currently held. + + The transfer manager only tracks requests from their final chunk + onward, so the transceiver has to be asked about pipelined chunks that + are already in flight during context compute. + """ + if self.async_transfer_manager.has_any_inflight_requests(): + return True + return (self.kv_cache_transceiver is not None + and self.kv_cache_transceiver.has_any_inflight_transfer()) + @nvtx_range("_check_kv_transfer_timeout") def _check_kv_transfer_timeout(self): if not self.kv_cache_transceiver: @@ -6012,6 +6036,8 @@ def flag_if_kv_transfer_timed_out(req: LlmRequest, type: str) -> None: f"kv_transfer_timeout_ms={timeout_ms}ms") req.py_kv_transfer_timed_out = True + # Context requests start their clock on the last chunk, which is also when + # they enter the transfer manager, so this covers the whole context side. for req in self.async_transfer_manager.requests_in_transfer().values(): flag_if_kv_transfer_timed_out(req, "context") @@ -6526,23 +6552,47 @@ def kv_connector_request_finished(req: LlmRequest): self.async_transfer_manager.start_transfer(req) if self.kv_cache_transceiver: + # A pending cancel that could not complete (a chunk was mid-write) + # leaves the request active and still being prefilled. Its session + # is already CANCELLED, so further chunks would only produce + # spurious FAILED results on the receiver. + cancel_pending_ids = set(self.canceled_req_ids) for req in scheduled_requests: - if req.is_context_only_request and ( - req.is_context_finished or req.is_finished_due_to_length - ) and not req.is_finished_due_to_cancellation: - # Forward is done for this request — release the - # IndexMapper slot so new requests can reuse it. - # KV blocks stay allocated for the upcoming transfer. - if hasattr(self.kv_cache_manager, 'release_index_slot'): - self.kv_cache_manager.release_index_slot( - req.py_request_id) - # Order is important here: we need to start the transfer before responding - # to make sure the blocks are stored for reuse before they are sent. - self.async_transfer_manager.start_transfer(req) - self.kv_cache_transceiver.respond_and_send_async(req) - - if self.kv_cache_transceiver.kv_transfer_timeout_ms is not None: - req.py_kv_transfer_start_time = time.monotonic() + if req.is_context_only_request and not req.is_finished_due_to_cancellation: + if self.kv_cache_transceiver.has_retired_send_session(req): + # The peer registration went away with the session, so + # no further slice can land. Checked before the branch + # below because start_transfer() would otherwise pin + # blocks that only end_transfer() can release. + req.state = LlmRequestState.DISAGG_TRANS_ERROR + continue + if req.is_context_finished or req.is_finished_due_to_length: + # Forward is done for this request — release the + # IndexMapper slot so new requests can reuse it. + # KV blocks stay allocated for the upcoming transfer. + if hasattr(self.kv_cache_manager, 'release_index_slot'): + self.kv_cache_manager.release_index_slot( + req.py_request_id) + # Order matters: start_transfer commits the request's blocks to the reuse + # tree and pins them, and must run before respond_and_send_async sends the + # final KV slice and (for the Python transceiver) transitions the request toward completion. + self.async_transfer_manager.start_transfer(req) + + # send KV slice for monolithic transfer or last chunk of pipelined transfer + self.kv_cache_transceiver.respond_and_send_async(req) + + if self.kv_cache_transceiver.kv_transfer_timeout_ms is not None: + req.py_kv_transfer_start_time = time.monotonic() + elif (self.kv_cache_transceiver.pipeline_transfer_enabled + and req.state != LlmRequestState.GENERATION_COMPLETE + and + (req.py_request_id if not req.is_child else + req.parent_request_id) not in cancel_pending_ids): + # send intermediate chunk for pipelined transfer. + # GENERATION_COMPLETE means an error path already failed + # and freed this request; _update_request_states skips + # those, so its chunk bounds are unset. + self.kv_cache_transceiver.respond_and_send_async(req) if self.kv_connector_manager: if not self.disable_overlap_scheduler: @@ -7086,11 +7136,19 @@ def _do_terminate_request(self, request: LlmRequest): self.result_wait_queues.pop(request.py_request_id, None) def _is_request_in_transmission(self, request) -> bool: - """Check if a request is currently in transmission state.""" - return (request.state - == LlmRequestState.DISAGG_CONTEXT_TRANS_IN_PROGRESS + """Check if a request's KV cache may still be read by the fabric. + + The request state only tracks the compute phase. Under pipelined + transfer a chunk can be in flight while the request is still being + prefilled, so the transceiver's own ownership record has to be + consulted as well. + """ + if (request.state == LlmRequestState.DISAGG_CONTEXT_TRANS_IN_PROGRESS or request.state - == LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS) + == LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS): + return True + return (self.kv_cache_transceiver is not None + and self.kv_cache_transceiver.has_inflight_transfer(request)) def _try_cancel_request(self, request) -> bool: """Check if a request can be canceled and attempt cancellation if needed. diff --git a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py index 680d607efeb5..6c64791b7068 100644 --- a/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py +++ b/tensorrt_llm/_torch/pyexecutor/py_executor_creator.py @@ -44,6 +44,7 @@ from .connectors.kv_cache_connector import KvCacheConnectorManager from .dwdp import DwdpManager from .guided_decoder import CapturableGuidedDecoder, GuidedDecoder +from .kv_cache_transceiver import resolve_cache_transceiver_config from .model_engine import PyTorchModelEngine from .model_loader import ModelLoader, _construct_checkpoint_loader from .py_executor import PyExecutor @@ -693,6 +694,9 @@ def drafting_loop_wrapper(model): max_num_tokens = model_engine.max_num_tokens sparse_attention_config = model_engine.sparse_attention_config + # Resolve this before cache reuse and cache manager selection consume it. + resolve_cache_transceiver_config(cache_transceiver_config) + config = model_engine.model.model_config.pretrained_config max_num_seq_slots = getattr(model_engine, "max_num_seq_slots", max_batch_size * getattr(mapping, "pp_size", 1)) diff --git a/tensorrt_llm/commands/serve.py b/tensorrt_llm/commands/serve.py index e5fc8cdb9f75..23f1a526541b 100644 --- a/tensorrt_llm/commands/serve.py +++ b/tensorrt_llm/commands/serve.py @@ -1879,9 +1879,8 @@ def disaggregated( if "--config_file" in sys.argv: logger.warning("--config_file is deprecated, use --config instead.") - disagg_cfg = parse_disagg_config_file(config_file) - if schedule_style: - disagg_cfg.schedule_style = schedule_style + disagg_cfg = parse_disagg_config_file( + config_file, schedule_style_override=schedule_style) # Generate a shared deployment ID for all workers in this disagg deployment. # Inherited by child processes via env var; used for deduplication at query time. diff --git a/tensorrt_llm/llmapi/disagg_utils.py b/tensorrt_llm/llmapi/disagg_utils.py index b498fbd85635..6314ab06b8ef 100644 --- a/tensorrt_llm/llmapi/disagg_utils.py +++ b/tensorrt_llm/llmapi/disagg_utils.py @@ -5,6 +5,7 @@ import uuid from dataclasses import dataclass, field from enum import IntEnum +from os import PathLike from typing import Any, Dict, List, Literal, Optional, Tuple import yaml @@ -74,6 +75,38 @@ def _extract_internal_request_auth_key( return top_level_key +def _validate_disagg_config(schedule_style: str, context_servers: dict, + generation_servers: dict) -> None: + valid_schedule_styles = {"context_first", "generation_first"} + if schedule_style not in valid_schedule_styles: + raise ValueError( + f"schedule_style must be one of {sorted(valid_schedule_styles)}, " + f"got {schedule_style!r}") + + for server_group, server_config in ( + ("context_servers", context_servers), + ("generation_servers", generation_servers), + ): + cache_transceiver_config = server_config.get("cache_transceiver_config") + if cache_transceiver_config is None: + continue + if not isinstance(cache_transceiver_config, dict): + raise ValueError( + f"{server_group}.cache_transceiver_config must be a mapping, " + f"got {type(cache_transceiver_config).__name__}") + if "enable_pipelined_transfer" not in cache_transceiver_config: + continue + enable_pipelined_transfer = validate_config_bool( + cache_transceiver_config["enable_pipelined_transfer"], + f"{server_group}.cache_transceiver_config.enable_pipelined_transfer" + ) + if (enable_pipelined_transfer and schedule_style != "generation_first"): + raise ValueError( + f"{server_group}.cache_transceiver_config." + "enable_pipelined_transfer=True requires top-level " + "schedule_style='generation_first'.") + + class ServerRole(IntEnum): CONTEXT = 0 GENERATION = 1 @@ -225,11 +258,18 @@ def get_ctx_gen_server_addrs( return ctx_server_urls, gen_server_urls -def parse_disagg_config_file(yaml_config_file: str): +def parse_disagg_config_file(yaml_config_file: str | PathLike[str], + schedule_style_override: Optional[str] = None): with open(yaml_config_file, 'r') as file: config = yaml.safe_load(file) + if config is None: + raise ValueError( + f"Disaggregated config file is empty: {yaml_config_file}") + if schedule_style_override is not None: + config["schedule_style"] = schedule_style_override + disagg_server_config = extract_disagg_cfg(**config) return disagg_server_config @@ -281,6 +321,8 @@ def extract_disagg_cfg(hostname: str = 'localhost', # Inherit the value from the top-level servers[key] = value + _validate_disagg_config(schedule_style, context_servers, generation_servers) + server_configs = [] disagg_cluster_config = None ctx_router_config = extract_router_config(context_servers) @@ -318,8 +360,7 @@ def extract_disagg_cfg(hostname: str = 'localhost', raise ValueError( f"node_id must be in range [0, {node_id_space}), got {node_id}") config.node_id = node_id - if schedule_style: - config.schedule_style = schedule_style + config.schedule_style = schedule_style config.allow_request_chat_template = validate_config_bool( allow_request_chat_template, "allow_request_chat_template") config.gen_strip_message_history = gen_strip_message_history diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 1afb2cc25994..679afcd0139b 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -4250,6 +4250,13 @@ class CacheTransceiverConfig(StrictBaseModel, PybindMirror): "Per-region size in MiB of the native-disagg KV-cache bounce buffer (one for send, one for recv). Bounce coalesces a request's scattered per-block KV into one contiguous fabric-VMM buffer and issues a single multi-rail NIXL write. The size doubles as the on/off switch: 0 (default) keeps the per-block path, >0 enables bounce at that capacity. Only used by the Python (v2) transceiver." ) + enable_pipelined_transfer: bool = Field( + default=False, + description="When True, start transferring each prefill chunk's KV cache " + "as soon as its prefill completes, overlapping GPU compute " + "with cache transfer. Requires enable_chunked_prefill=True and " + "schedule_style=generation_first and pipeline_parallel_size=1.") + def _resolve_default_backend(self) -> Tuple[Optional[str], Optional[str]]: """Effective backend after resolving "DEFAULT" against legacy env vars. @@ -4265,6 +4272,7 @@ def _resolve_default_backend(self) -> Tuple[Optional[str], Optional[str]]: return "NIXL", None def _to_pybind(self): + # enable_pipelined_transfer is consumed by the Python transceiver only and has no C++ counterpart. return _CacheTransceiverConfig( backend=_CacheTransceiverBackendType.from_string(self.backend), max_tokens_in_buffer=self.max_tokens_in_buffer, diff --git a/tensorrt_llm/usage/llm_args_golden_manifest.json b/tensorrt_llm/usage/llm_args_golden_manifest.json index 38436f22514a..73d64279d26a 100644 --- a/tensorrt_llm/usage/llm_args_golden_manifest.json +++ b/tensorrt_llm/usage/llm_args_golden_manifest.json @@ -159,6 +159,13 @@ "kind": "categorical", "path": "cache_transceiver_config.backend" }, + { + "allowed_values": [], + "annotation": "", + "converter": "", + "kind": "value", + "path": "cache_transceiver_config.enable_pipelined_transfer" + }, { "allowed_values": [], "annotation": "", diff --git a/tests/integration/defs/accuracy/test_disaggregated_serving.py b/tests/integration/defs/accuracy/test_disaggregated_serving.py index 749ed07f50a0..25814aaba928 100644 --- a/tests/integration/defs/accuracy/test_disaggregated_serving.py +++ b/tests/integration/defs/accuracy/test_disaggregated_serving.py @@ -738,6 +738,52 @@ def test_kv_cache_v2_nixl_python(self): self.MODEL_PATH) as llm: run_accuracy_test(llm, self.MODEL_NAME, ["GSM8K"]) + @skip_pre_hopper + @pytest.mark.skip_less_device(2) + @parametrize_with_ids("enable_block_reuse", [True, False]) + @parametrize_with_ids("disable_overlap_scheduler", [True, False]) + def test_pipelined_kv_transfer_nixl_python_accuracy( + self, enable_block_reuse: bool, disable_overlap_scheduler: bool): + """Test pipelined KV transfer accuracy using Python transceiver and C++ KVCacheManager.""" + kv_cache_config = { + "use_kv_cache_manager_v2": False, + "enable_block_reuse": enable_block_reuse, + } + cache_transceiver_config = { + "backend": "NIXL", + "transceiver_runtime": "PYTHON", + "max_tokens_in_buffer": 4096, + "enable_pipelined_transfer": True, + } + ctx_server_config = { + "max_num_tokens": 256, # cap prefill chunk size + "disable_overlap_scheduler": disable_overlap_scheduler, + "kv_cache_config": dict(kv_cache_config), + "cache_transceiver_config": dict(cache_transceiver_config), + "enable_chunked_prefill": True, + } + gen_server_config = { + "disable_overlap_scheduler": disable_overlap_scheduler, + "kv_cache_config": dict(kv_cache_config), + "cache_transceiver_config": dict(cache_transceiver_config), + "enable_chunked_prefill": True, + } + disaggregated_server_config = { + "hostname": "localhost", + "backend": "pytorch", + "schedule_style": "generation_first", + "context_servers": { + "num_instances": 1, + }, + "generation_servers": { + "num_instances": 1, + }, + } + with launch_disaggregated_llm(disaggregated_server_config, + ctx_server_config, gen_server_config, + self.MODEL_PATH) as llm: + run_accuracy_test(llm, self.MODEL_NAME, ["GSM8K"]) + @pytest.mark.skip_less_device(2) def test_ngram(self): speculative_decoding_config = { @@ -1544,6 +1590,55 @@ def test_kv_cache_v2_nixl_python(self, use_kv_cache_manager_v2): self.MODEL_PATH) as llm: run_accuracy_test(llm, self.MODEL_NAME, ["MMLU", "GSM8K"]) + @pytest.mark.skip_less_device(2) + @parametrize_with_ids("enable_block_reuse", [True]) + @parametrize_with_ids("disable_overlap_scheduler", [False]) + def test_pipelined_kv_transfer_nixl_python_accuracy( + self, enable_block_reuse: bool, disable_overlap_scheduler: bool): + """Test Python transceiver pipelined KV transfer accuracy for a VSWA model, Gemma 3.""" + kv_cache_config = { + "use_kv_cache_manager_v2": False, + "enable_block_reuse": enable_block_reuse, + "enable_partial_reuse": enable_block_reuse, + "max_attention_window": [512, 512, 512, 512, 512, 32768], + } + cache_transceiver_config = { + "backend": "NIXL", + "transceiver_runtime": "PYTHON", + "max_tokens_in_buffer": 4096, + "enable_pipelined_transfer": True, + } + ctx_server_config = { + "max_num_tokens": 256, + "disable_overlap_scheduler": disable_overlap_scheduler, + "cuda_graph_config": None, + "kv_cache_config": dict(kv_cache_config), + "cache_transceiver_config": dict(cache_transceiver_config), + "enable_chunked_prefill": True, + } + gen_server_config = { + "disable_overlap_scheduler": disable_overlap_scheduler, + "cuda_graph_config": None, + "kv_cache_config": dict(kv_cache_config), + "cache_transceiver_config": dict(cache_transceiver_config), + "enable_chunked_prefill": True, + } + disaggregated_server_config = { + "hostname": "localhost", + "backend": "pytorch", + "schedule_style": "generation_first", + "context_servers": { + "num_instances": 1, + }, + "generation_servers": { + "num_instances": 1, + }, + } + with launch_disaggregated_llm(disaggregated_server_config, + ctx_server_config, gen_server_config, + self.MODEL_PATH) as llm: + run_accuracy_test(llm, self.MODEL_NAME, ["GSM8K"]) + @skip_pre_blackwell @pytest.mark.skip_less_device_memory(80000) diff --git a/tests/integration/test_lists/test-db/l0_dgx_b200.yml b/tests/integration/test_lists/test-db/l0_dgx_b200.yml index 74d7936cab15..cda33bc3f700 100644 --- a/tests/integration/test_lists/test-db/l0_dgx_b200.yml +++ b/tests/integration/test_lists/test-db/l0_dgx_b200.yml @@ -18,6 +18,12 @@ l0_dgx_b200: - unittest/_torch/misc/test_autotuner.py::test_autotuner_distributed_strategy - accuracy/test_llm_api_pytorch.py::TestQwen3_5_35B_A3B::test_bf16[tp2-TRTLLM] - accuracy/test_disaggregated_serving.py::TestDeepSeekV3Lite::test_gen_only_sync[cpp] + # ------------- Disaggregated Serving: Pipelined KV Transfer (multi-GPU) --------------- + - accuracy/test_disaggregated_serving.py::TestLlama3_1_8BInstruct::test_pipelined_kv_transfer_nixl_python_accuracy[disable_overlap_scheduler=False-enable_block_reuse=False] + - accuracy/test_disaggregated_serving.py::TestLlama3_1_8BInstruct::test_pipelined_kv_transfer_nixl_python_accuracy[disable_overlap_scheduler=False-enable_block_reuse=True] + - accuracy/test_disaggregated_serving.py::TestLlama3_1_8BInstruct::test_pipelined_kv_transfer_nixl_python_accuracy[disable_overlap_scheduler=True-enable_block_reuse=False] + - accuracy/test_disaggregated_serving.py::TestLlama3_1_8BInstruct::test_pipelined_kv_transfer_nixl_python_accuracy[disable_overlap_scheduler=True-enable_block_reuse=True] + - accuracy/test_disaggregated_serving.py::TestGemma3_1BInstruct::test_pipelined_kv_transfer_nixl_python_accuracy[disable_overlap_scheduler=False-enable_block_reuse=True] # ------------- KV Cache V2 Scheduler IT (multi-GPU) --------------- - kv_cache/test_kv_cache_v2_scheduler.py::TestKVCacheV2DSv3Lite::test_mtp_draft_tokens - kv_cache/test_kv_cache_v2_scheduler.py::TestKVCacheV2DSv3Lite::test_mtp_chunked_draft_tokens diff --git a/tests/unittest/_torch/executor/test_disagg_index_mapper_early_release.py b/tests/unittest/_torch/executor/test_disagg_index_mapper_early_release.py index c6d1a236ff88..a8c448d26224 100644 --- a/tests/unittest/_torch/executor/test_disagg_index_mapper_early_release.py +++ b/tests/unittest/_torch/executor/test_disagg_index_mapper_early_release.py @@ -77,6 +77,7 @@ def __init__(self, kv_cache_manager, async_transfer_manager, kv_cache_transceive self.kv_connector_manager = None self.disable_overlap_scheduler = True self.previous_batch = None + self.canceled_req_ids = [] def _check_disagg_ctx_cache_transfer_status(self, _): return None @@ -92,6 +93,7 @@ def _build(self, kv_cache_manager): transfer_manager = AsyncTransferManager(resource_manager) transceiver = MagicMock() transceiver.kv_transfer_timeout_ms = None + transceiver.has_retired_send_session.return_value = False return _FakeExecutor(kv_cache_manager, transfer_manager, transceiver), transfer_manager def test_send_kv_async_calls_release_index_slot(self): diff --git a/tests/unittest/_torch/executor/test_disagg_inflight_cancel_gate.py b/tests/unittest/_torch/executor/test_disagg_inflight_cancel_gate.py index 08a6f4701761..c892352a59e9 100644 --- a/tests/unittest/_torch/executor/test_disagg_inflight_cancel_gate.py +++ b/tests/unittest/_torch/executor/test_disagg_inflight_cancel_gate.py @@ -300,6 +300,8 @@ def test_user_cancel_waits_for_context_transfer_owners(monkeypatch): executor.kv_cache_transceiver = Mock() executor.kv_cache_transceiver.cancel_request.return_value = True executor.kv_cache_transceiver.supports_inflight_request_cancellation.return_value = True + # Transfer ownership is released once the request leaves the manager. + executor.kv_cache_transceiver.has_inflight_transfer.return_value = False executor._disagg_inflight_cancel_unsupported_logged = False executor.async_transfer_manager = Mock() executor.async_transfer_manager.requests_in_transfer.return_value = { diff --git a/tests/unittest/disaggregated/test_bounce.py b/tests/unittest/disaggregated/test_bounce.py index 84a7eda4fe1a..00132e5d27bd 100644 --- a/tests/unittest/disaggregated/test_bounce.py +++ b/tests/unittest/disaggregated/test_bounce.py @@ -194,16 +194,25 @@ def test_encode_tail_handles_unset_base(self): def test_kv_result_prefix_roundtrip(): """The KV_AGENT_RESULT binary prefix (transfer.py) must round-trip exactly.""" tfr = pytest.importorskip("tensorrt_llm._torch.disaggregation.native.transfer") - for rank, rid, sl, last, status, size in [ - (7, 6925227277844486, 42, True, tfr.AgentResult.SUCCESS, 4096), - (0, 1, 0, False, tfr.AgentResult.FAILED, 0), - (31, 2**62, 9999, True, tfr.AgentResult.SUCCESS, 2**40), + # Sender and receiver slice ids differ in every case so a field swap is caught. + for rank, rid, send_sl, recv_sl, last, status, size in [ + (7, 6925227277844486, 42, 0, True, tfr.AgentResult.SUCCESS, 4096), + (0, 1, 0, 3, False, tfr.AgentResult.FAILED, 0), + (31, 2**62, 9999, 1, True, tfr.AgentResult.SUCCESS, 2**40), + (2, 77, tfr.NO_SLICE_ID, 0, True, tfr.AgentResult.FAILED, 0), ]: packed = tfr._KV_RESULT_PREFIX.pack( - rank, rid, sl, last, tfr._AGENT_RESULT_CODE[status], size + rank, rid, send_sl, recv_sl, last, tfr._AGENT_RESULT_CODE[status], size + ) + r, i, send_out, recv_out, last_out, c, sz = tfr._KV_RESULT_PREFIX.unpack(packed) + assert (r, i, send_out, recv_out, last_out, sz) == ( + rank, + rid, + send_sl, + recv_sl, + last, + size, ) - r, i, s, last_out, c, sz = tfr._KV_RESULT_PREFIX.unpack(packed) - assert (r, i, s, last_out, sz) == (rank, rid, sl, last, size) assert tfr._AGENT_RESULT_BY_CODE[c] is status @@ -212,11 +221,11 @@ def test_make_kv_result_msg_uses_binary_frame(result_name): """Every KV result (success and failure) uses the binary frame so the receiver can decode it.""" tfr = pytest.importorskip("tensorrt_llm._torch.disaggregation.native.transfer") result = getattr(tfr.AgentResult, result_name) - msg = tfr._make_kv_result_msg(3, 12345, 7, True, result, transfer_size=8192) + msg = tfr._make_kv_result_msg(3, 12345, 7, 2, True, result, transfer_size=8192) assert msg[0] == tfr.MessageType.KV_AGENT_RESULT assert len(msg) == 2 # prefix only; no bounce tail when none is passed - r, rid, sl, last, code, size = tfr._KV_RESULT_PREFIX.unpack(msg[1]) - assert (r, rid, sl, last, size) == (3, 12345, 7, True, 8192) + r, rid, send_sl, recv_sl, last, code, size = tfr._KV_RESULT_PREFIX.unpack(msg[1]) + assert (r, rid, send_sl, recv_sl, last, size) == (3, 12345, 7, 2, True, 8192) assert tfr._AGENT_RESULT_BY_CODE[code] is result diff --git a/tests/unittest/disaggregated/test_chunked_transfer.py b/tests/unittest/disaggregated/test_chunked_transfer.py new file mode 100644 index 000000000000..43362f75bab0 --- /dev/null +++ b/tests/unittest/disaggregated/test_chunked_transfer.py @@ -0,0 +1,1403 @@ +# 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. +"""Unit tests for chunked and pipelined KV cache transfer (sender-only chunking). + +These tests validate the session state machine using the real +TxSession/RxSession classes with lightweight stub sender/receiver objects. +""" + +import threading +import time +from types import MethodType, SimpleNamespace +from unittest.mock import MagicMock + +import numpy as np +import pytest + +from tensorrt_llm import DisaggregatedParams +from tensorrt_llm._torch.disaggregation.base.transfer import ( + KVSlice, + SessionStatus, + TokenRange, + WaitResult, + project_blocks_to_global_chunk, +) +from tensorrt_llm._torch.disaggregation.native.transfer import ( + _KV_RESULT_PREFIX, + NO_SLICE_ID, + AgentResult, + KVSendTask, + RecvReqInfo, + RxSession, + Sender, + TaskStatus, + TxSession, + WriteMeta, +) +from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequestState +from tensorrt_llm.disaggregated_params import DisaggScheduleStyle +from tensorrt_llm.llmapi.llm_args import CacheTransceiverConfig + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_params(rid: int = 42) -> DisaggregatedParams: + return DisaggregatedParams(disagg_request_id=rid) + + +def _stub_sender(): + """Create a stub sender with no-op methods needed by TxSession.""" + sender = MagicMock() + sender.setup_session = MagicMock() + sender._get_req_info = MagicMock(return_value=None) + sender.dispatch_task = MagicMock() + return sender + + +def _stub_receiver(): + """Create a stub receiver with no-op methods needed by RxSession.""" + receiver = MagicMock() + receiver.setup_session = MagicMock() + receiver.dispatch_task = MagicMock() + return receiver + + +def _make_tx_session(num_slices: int, rid: int = 42, prompt_len: int = 8, **kwargs) -> TxSession: + """Create a real TxSession and send num_slices slices into it.""" + params = _make_params(rid) + session = TxSession( + request_id=rid, + params=params, + sender=_stub_sender(), + prompt_len=prompt_len, + **kwargs, + ) + for i in range(num_slices): + s = KVSlice( + is_last_slice=(i == num_slices - 1), + block_ids_per_layer_groups=[[i]], + ) + session.send(s) + return session + + +def _make_rx_session(num_slices: int, rid: int = 42, prompt_len: int = 8) -> RxSession: + """Create a real RxSession and receive num_slices slices into it.""" + params = _make_params(rid) + session = RxSession( + request_id=rid, + params=params, + receiver=_stub_receiver(), + prompt_len=prompt_len, + ) + for i in range(num_slices): + s = KVSlice( + is_last_slice=(i == num_slices - 1), + block_ids_per_layer_groups=[[i]], + ) + session.receive(s) + return session + + +def _make_replay_sender(rid: int = 42) -> tuple[Sender, TxSession]: + """Create a Sender/TxSession pair with controllable replay dispatch.""" + sender = Sender.__new__(Sender) + sender._sessions_lock = threading.Lock() + sender._req_infos = {} + sender._session = None + sender.setup_session = lambda session: setattr(sender, "_session", session) + sender._get_session = lambda unique_rid: sender._session if unique_rid == rid else None + sender._get_req_info = lambda unique_rid: sender._req_infos.get(unique_rid) + + def save_peer_req_info(info): + sender._req_infos.setdefault(info.unique_rid, {})[info.instance_rank] = info + + sender._save_peer_req_info = save_peer_req_info + sender._send_failed_result_to_receiver = MagicMock() + sender.dispatch_task = MethodType(Sender.dispatch_task, sender) + + session = TxSession( + request_id=rid, + params=_make_params(rid), + sender=sender, + prompt_len=8, + ) + return sender, session + + +def _replay_info(rid: int = 42) -> RecvReqInfo: + return RecvReqInfo( + sender_req_id=rid, + instance_name="decode", + instance_rank=0, + block_ids_per_layer_groups=[], + unique_rid=rid, + ) + + +# --------------------------------------------------------------------------- +# Global chunk projection tests +# --------------------------------------------------------------------------- + + +def test_chunk_projection_noops_when_chunk_is_outside_short_layer_group(): + """A shared chunk cursor past a short layer group's resident range is a no-op.""" + block_ids = np.array([10, 11, 12], dtype=np.int64) + + projected_ids = project_blocks_to_global_chunk( + block_ids, + chunk_block_offset=4, + chunk_block_count=4, + resident_block_end=3, + ) + + assert projected_ids.size == 0 + + +@pytest.mark.parametrize( + "resident_block_end,chunk_block_offset,expected", + [ + (16, 0, np.arange(16, dtype=np.int64)), + (32, 16, np.arange(16, 32, dtype=np.int64)), + ], + ids=["first_chunk", "later_chunk"], +) +def test_chunk_projection_maps_incrementally_allocated_source( + resident_block_end, chunk_block_offset, expected +): + """Source blocks end at the current chunk, not at the full prompt.""" + block_ids = np.arange(resident_block_end, dtype=np.int64) + + projected_ids = project_blocks_to_global_chunk( + block_ids, + chunk_block_offset=chunk_block_offset, + chunk_block_count=16, + resident_block_end=resident_block_end, + ) + + assert np.array_equal(projected_ids, expected) + + +def test_chunk_projection_maps_prefix_reuse_suffix_by_overlap(): + """Destination suffixes are matched by overlap, not by raw chunk-offset indexing.""" + block_ids = np.array([104, 105, 106, 107], dtype=np.int64) + + first_chunk = project_blocks_to_global_chunk( + block_ids, + chunk_block_offset=0, + chunk_block_count=4, + resident_block_end=8, + ) + second_chunk = project_blocks_to_global_chunk( + block_ids, + chunk_block_offset=4, + chunk_block_count=4, + resident_block_end=8, + ) + + assert first_chunk.size == 0 + assert np.array_equal(second_chunk, block_ids) + + +def _make_projection_sender() -> Sender: + """Create a Sender wired to a stub registrar with two non-windowed layer groups.""" + peer_ri = SimpleNamespace( + dp_rank=0, + device_id=0, + instance_name="decode", + instance_rank=0, + self_endpoint="tcp://decode:0", + ) + + extractor = MagicMock() + extractor.page_table = SimpleNamespace( + tokens_per_block=8, + layer_groups=[ + SimpleNamespace(sliding_window_size=None), + SimpleNamespace(sliding_window_size=None), + ], + ) + extractor.extract.side_effect = lambda block_ids, **_: SimpleNamespace( + memory=SimpleNamespace( + ptrs=np.asarray(block_ids, dtype=np.int64), + bytes_per_region=1, + ) + ) + + mapper = MagicMock() + mapper.map.side_effect = lambda src_region, dst_region: SimpleNamespace( + src=src_region, + dst=dst_region, + ) + + registrar = MagicMock() + registrar.self_rank_info = SimpleNamespace() + registrar.self_extractor = extractor + registrar.get_peer_rank_info.return_value = peer_ri + registrar.get_peer_overlap.return_value = SimpleNamespace(ranks=[0]) + registrar.should_send_kv.return_value = True + registrar.get_pool_mapping.return_value = { + (0, 0): (0, 0), + (1, 0): (1, 0), + } + registrar.peer_extractor.return_value = extractor + registrar.get_kv_map.return_value = mapper + + sender = Sender.__new__(Sender) + sender._registrar = registrar + return sender + + +def _make_projection_task(slice_id: int = 1) -> KVSendTask: + return KVSendTask( + KVSlice( + is_last_slice=True, + block_ids_per_layer_groups=[ + np.array([4, 5, 6, 7], dtype=np.int64), + np.array([10, 11, 12], dtype=np.int64), + ], + token_range=TokenRange(start=32, end=64), + total_blocks=8, + ), + _make_params(), + slice_id=slice_id, + prompt_len=64, + ) + + +def _make_projection_req_info(slice_id=None) -> RecvReqInfo: + return RecvReqInfo( + sender_req_id=42, + instance_name="decode", + instance_rank=0, + block_ids_per_layer_groups=[ + np.array([104, 105, 106, 107], dtype=np.int64), + np.array([200, 201, 202], dtype=np.int64), + ], + unique_rid=42, + slice_id=slice_id, + ) + + +def test_build_kv_write_meta_projects_asymmetric_layer_group_chunk(): + """A short layer group's suffix blocks transfer with the overlapping global chunk.""" + sender = _make_projection_sender() + + write_meta = sender._build_kv_write_meta(_make_projection_task(), _make_projection_req_info()) + + assert np.array_equal( + write_meta.src_ptrs, + np.array([4, 5, 6, 7, 10, 11, 12], dtype=np.int64), + ) + assert np.array_equal( + write_meta.dst_ptrs, + np.array([104, 105, 106, 107, 200, 201, 202], dtype=np.int64), + ) + assert np.array_equal(write_meta.sizes, np.ones(7, dtype=np.int64)) + # The sender's chunk index and the peer's task index are tracked separately; + # a receiver that sends no slice_id is addressed as its single task 0. + assert write_meta.sender_slice_id == 1 + assert write_meta.receiver_slice_id == 0 + + +def test_build_kv_write_meta_echoes_receiver_slice_id(): + """receiver_slice_id comes from the peer's RecvReqInfo, not the sender's chunk index.""" + sender = _make_projection_sender() + + write_meta = sender._build_kv_write_meta( + _make_projection_task(slice_id=1), _make_projection_req_info(slice_id=3) + ) + + assert write_meta.sender_slice_id == 1 + assert write_meta.receiver_slice_id == 3 + + +# --------------------------------------------------------------------------- +# KV_AGENT_RESULT slice-id addressing tests +# --------------------------------------------------------------------------- + + +def _make_write_meta(sender_slice_id, receiver_slice_id) -> WriteMeta: + empty = np.array([], dtype=np.int64) + return WriteMeta( + task=MagicMock(), + expected_transfers=1, + peer_name="decode0", + peer_rank=0, + peer_endpoint="tcp://decode:0", + unique_rid=42, + src_ptrs=empty, + dst_ptrs=empty, + sizes=empty, + sender_slice_id=sender_slice_id, + receiver_slice_id=receiver_slice_id, + ) + + +def test_send_kv_result_carries_both_slice_ids(): + """The result frame reports the sender's chunk and addresses the peer's own task.""" + sender = Sender.__new__(Sender) + sender._instance_rank = 5 + dealer = MagicMock() + sender._get_or_connect_thread_dealer = MagicMock(return_value=dealer) + + sender._send_kv_result_to_receiver( + _make_write_meta(sender_slice_id=4, receiver_slice_id=2), + is_last=True, + result=AgentResult.SUCCESS, + ) + + (msg,), _ = dealer.send.call_args + (rank, rid, sender_slice_id, receiver_slice_id, is_last, _code, _size) = ( + _KV_RESULT_PREFIX.unpack(msg[1]) + ) + assert (rank, rid, sender_slice_id, receiver_slice_id, is_last) == (5, 42, 4, 2, True) + + +def test_send_kv_result_without_task_reports_no_slice_id(): + """A result with no owning KVSendTask reports NO_SLICE_ID rather than chunk 0.""" + sender = Sender.__new__(Sender) + sender._instance_rank = 5 + dealer = MagicMock() + sender._get_or_connect_thread_dealer = MagicMock(return_value=dealer) + + sender._send_kv_result_to_receiver( + _make_write_meta(sender_slice_id=None, receiver_slice_id=0), + is_last=True, + result=AgentResult.FAILED, + ) + + (msg,), _ = dealer.send.call_args + (_rank, _rid, sender_slice_id, _receiver_slice_id, _is_last, _code, _size) = ( + _KV_RESULT_PREFIX.unpack(msg[1]) + ) + assert sender_slice_id == NO_SLICE_ID + + +def test_process_kv_agent_result_resolves_task_by_receiver_slice_id(): + """A sender chunk id far past the receiver's task count still resolves task 0.""" + session = _make_rx_session(1) + session._kv_tasks[0].expected_transfers = 1 + session._receiver._bounce.is_bounced.return_value = False + + session.process_kv_agent_result( + peer_rank=0, + receiver_slice_id=0, + sender_slice_id=4, + is_last_slice=True, + status=AgentResult.SUCCESS, + ) + + assert session._kv_tasks[0].status == TaskStatus.TRANSFERRED + + +def test_process_kv_agent_result_rejects_unknown_receiver_slice_id(): + """Indexing is bounded by the receiver's own task count, and names both ids.""" + session = _make_rx_session(1) + + with pytest.raises(AssertionError, match=r"receiver_slice_id=2.*sender_slice_id=0"): + session.process_kv_agent_result( + peer_rank=0, + receiver_slice_id=2, + sender_slice_id=0, + is_last_slice=True, + status=AgentResult.SUCCESS, + ) + + +def test_process_kv_agent_result_failure_attributes_sender_chunk(): + """A failed chunk names the sender chunk so the failure is attributable.""" + session = _make_rx_session(1) + + session.process_kv_agent_result( + peer_rank=0, + receiver_slice_id=0, + sender_slice_id=3, + is_last_slice=False, + status=AgentResult.FAILED, + ) + + task = session._kv_tasks[0] + assert task.status == TaskStatus.ERROR + assert "sender_slice_id=3" in str(task._exception) + assert session.status == SessionStatus.ERROR + + +# --------------------------------------------------------------------------- +# TxSession multi-slice status tests (real class) +# --------------------------------------------------------------------------- + + +def test_late_peer_replay_enqueues_before_concurrent_final_slice(): + """Replay holds dispatch ordering until every older slice is queued.""" + sender, session = _make_replay_sender() + enqueued_slice_ids = [] + replay_build_started = threading.Event() + release_replay = threading.Event() + final_send_started = threading.Event() + final_send_done = threading.Event() + + def build_write_meta(task, info): + if task.slice_id == 0 and not replay_build_started.is_set(): + replay_build_started.set() + assert release_replay.wait(timeout=5) + return SimpleNamespace(task=task, peer_rank=info.instance_rank) + + sender._build_kv_write_meta = build_write_meta + sender._enqueue = lambda meta: enqueued_slice_ids.append(meta.task.slice_id) + + session.send(KVSlice(is_last_slice=False, block_ids_per_layer_groups=[[0]])) + session.send(KVSlice(is_last_slice=False, block_ids_per_layer_groups=[[1]])) + + info = _replay_info() + replay_thread = threading.Thread( + target=sender._respond_with_kv, + args=(b"", [b"REQUEST_DATA", info.to_bytes()]), + ) + replay_thread.start() + assert replay_build_started.wait(timeout=5) + + def send_final_slice(): + final_send_started.set() + session.send(KVSlice(is_last_slice=True, block_ids_per_layer_groups=[[2]])) + final_send_done.set() + + final_thread = threading.Thread(target=send_final_slice) + final_thread.start() + assert final_send_started.wait(timeout=5) + assert not final_send_done.wait(timeout=0.1) + + release_replay.set() + replay_thread.join(timeout=5) + final_thread.join(timeout=5) + + assert not replay_thread.is_alive() + assert not final_thread.is_alive() + assert enqueued_slice_ids == [0, 1, 2] + + +def test_late_peer_replay_includes_final_slice_sent_before_registration(): + """A final slice buffered before replay is queued once after older slices.""" + sender, session = _make_replay_sender() + enqueued_slice_ids = [] + sender._build_kv_write_meta = lambda task, info: SimpleNamespace( + task=task, peer_rank=info.instance_rank + ) + sender._enqueue = lambda meta: enqueued_slice_ids.append(meta.task.slice_id) + + session.send(KVSlice(is_last_slice=False, block_ids_per_layer_groups=[[0]])) + session.send(KVSlice(is_last_slice=False, block_ids_per_layer_groups=[[1]])) + session.send(KVSlice(is_last_slice=True, block_ids_per_layer_groups=[[2]])) + + info = _replay_info() + sender._respond_with_kv(b"", [b"REQUEST_DATA", info.to_bytes()]) + + assert enqueued_slice_ids == [0, 1, 2] + + +def test_tx_session_status_init_until_all_transferred(): + """TxSession status is not KV_TRANSFERRED until ALL tasks complete.""" + session = _make_tx_session(3) + session.receiver_ready = True + assert session.status == SessionStatus.TRANSFERRING or session.status == SessionStatus.READY + + session.kv_tasks[0].status = TaskStatus.TRANSFERRED + assert session.status != SessionStatus.KV_TRANSFERRED + + session.kv_tasks[1].status = TaskStatus.TRANSFERRED + assert session.status != SessionStatus.KV_TRANSFERRED + + session.kv_tasks[2].status = TaskStatus.TRANSFERRED + assert session.status == SessionStatus.KV_TRANSFERRED + + +def test_tx_session_status_error_on_any_failure(): + """TxSession status is ERROR if any task fails.""" + session = _make_tx_session(3) + session.kv_tasks[0].status = TaskStatus.TRANSFERRED + session.kv_tasks[1].status = TaskStatus.ERROR + assert session.status == SessionStatus.ERROR + + +def test_tx_session_wait_complete_all_tasks(): + """TxSession.wait_complete blocks on all task futures.""" + session = _make_tx_session(3) + for task in session.kv_tasks: + task.complete() + + result = session.wait_complete() + assert result == WaitResult.COMPLETED + + +def test_tx_session_wait_complete_fails_on_partial_failure(): + """TxSession.wait_complete returns FAILED if any task fails.""" + session = _make_tx_session(3) + session.kv_tasks[0].complete() + session.kv_tasks[1].fail(RuntimeError("transfer failed")) + session.kv_tasks[2].complete() + + result = session.wait_complete() + assert result == WaitResult.FAILED + + +# --------------------------------------------------------------------------- +# RxSession multi-slice status tests (real class) +# --------------------------------------------------------------------------- + + +def test_rx_session_status_checks_all_tasks(): + """RxSession status is KV_TRANSFERRED only when ALL tasks complete.""" + session = _make_rx_session(3) + assert session.status == SessionStatus.INIT + + session._kv_tasks[0].status = TaskStatus.TRANSFERRED + session._kv_tasks[1].status = TaskStatus.TRANSFERRING + assert session.status == SessionStatus.TRANSFERRING + + session._kv_tasks[1].status = TaskStatus.TRANSFERRED + session._kv_tasks[2].status = TaskStatus.TRANSFERRED + assert session.status == SessionStatus.KV_TRANSFERRED + + +def test_rx_session_status_error_on_any_failure(): + """RxSession status is ERROR if any task fails.""" + session = _make_rx_session(2) + session._kv_tasks[0].status = TaskStatus.TRANSFERRED + session._kv_tasks[1].status = TaskStatus.ERROR + assert session.status == SessionStatus.ERROR + + +def test_rx_session_process_aux_completes_at_expected_transfers(): + """Aux completes only once the expected transfer count is reached. + + The receiver always has exactly one task. + """ + session = _make_rx_session(1) + session._kv_tasks[0].expected_transfers = 2 + + session.process_aux_agent_result(0, AgentResult.SUCCESS) + assert session._aux_status != TaskStatus.TRANSFERRED + + session.process_aux_agent_result(0, AgentResult.SUCCESS) + assert session._aux_status == TaskStatus.TRANSFERRED + + +def test_rx_session_wait_complete_all_tasks(): + """RxSession.wait_complete blocks on all task futures.""" + session = _make_rx_session(3) + for task in session._kv_tasks: + task.complete() + + result = session.wait_complete() + assert result == WaitResult.COMPLETED + + +def test_rx_session_wait_complete_fails_on_partial_failure(): + """RxSession.wait_complete returns FAILED if any task fails.""" + session = _make_rx_session(2) + session._kv_tasks[0].complete() + session._kv_tasks[1].fail(RuntimeError("transfer failed")) + + result = session.wait_complete() + assert result == WaitResult.FAILED + + +# --------------------------------------------------------------------------- +# Mid-transfer chunk failure tests +# --------------------------------------------------------------------------- + + +def test_tx_session_mid_chunk_failure(): + """If one chunk fails mid-transfer, the session reports ERROR.""" + session = _make_tx_session(4) + + session.kv_tasks[0].complete() + session.kv_tasks[1].complete() + session.kv_tasks[2].fail(RuntimeError("RDMA failed")) + session.kv_tasks[3].complete() + + assert session.status == SessionStatus.ERROR + result = session.wait_complete() + assert result == WaitResult.FAILED + + +def test_rx_session_mid_chunk_failure(): + """If one chunk fails mid-transfer on receiver, the session reports ERROR.""" + session = _make_rx_session(4) + + session._kv_tasks[0].complete() + session._kv_tasks[1].fail(RuntimeError("RDMA failed")) + session._kv_tasks[2].complete() + session._kv_tasks[3].complete() + + assert session.status == SessionStatus.ERROR + result = session.wait_complete() + assert result == WaitResult.FAILED + + +# --------------------------------------------------------------------------- +# Pipelined transfer tests +# --------------------------------------------------------------------------- + + +def test_pipelined_transfer_disabled_by_default(): + """pipeline_transfer_enabled reflects the configured flag.""" + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 + + transceiver = MagicMock() + transceiver._enable_pipelined_transfer = False + transceiver._chunk_size_blocks = 64 + + result = KvCacheTransceiverV2.pipeline_transfer_enabled.fget(transceiver) + assert result is False + + +def test_pipelined_transfer_requires_chunked_prefill(): + """ValueError when pipelined transfer is enabled without chunked prefill.""" + from tensorrt_llm._torch.pyexecutor.kv_cache_transceiver import create_kv_cache_transceiver + + cache_transceiver_config = CacheTransceiverConfig( + backend="NIXL", + enable_pipelined_transfer=True, + ) + + with pytest.raises( + ValueError, + match="enable_chunked_prefill is required when enable_pipelined_transfer is set.", + ): + create_kv_cache_transceiver( + MagicMock(), + MagicMock(), + MagicMock(), + MagicMock(), + cache_transceiver_config, + enable_chunked_prefill=False, + ) + + +def test_pipelined_transfer_rejects_pipeline_parallelism(monkeypatch): + """ValueError for pipeline parallelism when the disaggregated role is unknown.""" + from tensorrt_llm._torch.pyexecutor.kv_cache_transceiver import create_kv_cache_transceiver + + monkeypatch.delenv("TRTLLM_DISAGG_ROLE", raising=False) + mapping = MagicMock() + mapping.pp_size = 2 + cache_transceiver_config = CacheTransceiverConfig( + backend="NIXL", + enable_pipelined_transfer=True, + ) + + with pytest.raises( + ValueError, + match="pipeline_parallel_size=1 is required when enable_pipelined_transfer is set.", + ): + create_kv_cache_transceiver( + mapping, + MagicMock(), + MagicMock(), + MagicMock(), + cache_transceiver_config, + enable_chunked_prefill=True, + ) + + +def test_pipelined_transfer_allows_pipeline_parallelism_on_generation_server(monkeypatch): + """Pipeline parallelism is allowed when the worker only receives KV cache.""" + from tensorrt_llm._torch.disaggregation import transceiver as transceiver_module + from tensorrt_llm._torch.pyexecutor.kv_cache_transceiver import create_kv_cache_transceiver + + monkeypatch.setenv("TRTLLM_DISAGG_ROLE", "generation") + transceiver = MagicMock() + transceiver_cls = MagicMock(return_value=transceiver) + monkeypatch.setattr(transceiver_module, "KvCacheTransceiverV2", transceiver_cls) + + mapping = MagicMock() + mapping.pp_size = 2 + cache_transceiver_config = CacheTransceiverConfig( + backend="NIXL", + enable_pipelined_transfer=True, + ) + + result = create_kv_cache_transceiver( + mapping, + MagicMock(), + MagicMock(), + MagicMock(), + cache_transceiver_config, + enable_chunked_prefill=True, + ) + + assert result is transceiver + assert cache_transceiver_config.transceiver_runtime == "PYTHON" + transceiver_cls.assert_called_once() + + +def test_python_transceiver_rejects_cpp_mamba_cache_manager(): + """Python transceiver requires separate Python-managed Mamba state.""" + from tensorrt_llm._torch.pyexecutor.kv_cache_transceiver import create_kv_cache_transceiver + from tensorrt_llm._torch.pyexecutor.mamba_cache_manager import CppMambaHybridCacheManager + + kv_cache_manager = object.__new__(CppMambaHybridCacheManager) + cache_transceiver_config = CacheTransceiverConfig( + backend="NIXL", + transceiver_runtime="PYTHON", + ) + + # A hybrid manager arrives as both kv_cache_manager and mamba_cache_manager, + # the way _util.py passes it. + with pytest.raises( + ValueError, + match="cannot drive CppMambaHybridCacheManager", + ): + create_kv_cache_transceiver( + MagicMock(), + MagicMock(), + kv_cache_manager, + MagicMock(), + cache_transceiver_config, + kv_cache_manager, + ) + + +def test_pipelined_transfer_requires_gen_first_flow(): + """ValueError when a real request is not using gen-first flow.""" + from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor + + executor = MagicMock() + executor.is_warmup = False + executor.kv_cache_transceiver.pipeline_transfer_enabled = True + executor._validate_token_id_range = MagicMock() + executor.sampler.validate_request = MagicMock() + + request = MagicMock() + request.sampling_config = None + request.py_beam_width = 1 + request.py_disaggregated_params = SimpleNamespace( + schedule_style=DisaggScheduleStyle.CONTEXT_FIRST + ) + + with pytest.raises( + ValueError, + match="schedule_style must be generation_first when enable_pipelined_transfer is set.", + ): + PyExecutor._validate_request(executor, request) + + +def test_pipelined_transfer_allows_non_disaggregated_request(): + """Requests without disaggregated parameters do not transfer KV cache.""" + from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor + + executor = MagicMock() + executor.is_warmup = False + executor.max_beam_width = 1 + executor.kv_cache_transceiver.pipeline_transfer_enabled = True + executor._validate_token_id_range = MagicMock() + executor.sampler.validate_request = MagicMock() + + request = MagicMock() + request.sampling_config = None + request.py_beam_width = 1 + request.py_disaggregated_params = None + + PyExecutor._validate_request(executor, request) + + executor.sampler.validate_request.assert_called_once_with(request) + + +def test_pipelined_last_chunk_sends_and_finalizes(): + """respond_and_send_async sends the built chunk and finalizes on the last chunk.""" + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 + + session = MagicMock() + session.kv_tasks = [] + + last_slice = KVSlice( + is_last_slice=True, + block_ids_per_layer_groups=[np.array([0, 1], dtype=np.int64)], + ) + + transceiver = MagicMock() + transceiver._enable_pipelined_transfer = True + transceiver.kv_transfer_timeout_ms = None + transceiver._get_or_create_send_session.return_value = session + transceiver._build_prefill_chunk.return_value = last_slice + + request = SimpleNamespace( + py_disaggregated_params=DisaggregatedParams(disagg_request_id=42), + request_id=42, + prompt_len=8, + py_beam_width=1, + py_kv_transfer_start_time=None, + set_kv_cache_transfer_start=lambda _ts: None, + ) + + KvCacheTransceiverV2.respond_and_send_async(transceiver, request) + + transceiver._build_prefill_chunk.assert_called_once_with(request) + session.send.assert_called_once_with(last_slice) + transceiver._finalize_send.assert_called_once_with(request, session) + assert request.state == LlmRequestState.DISAGG_CONTEXT_TRANS_IN_PROGRESS + + +def test_pipelined_non_last_chunk_does_not_finalize(): + """respond_and_send_async sends non-final chunks without finalizing.""" + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 + + session = MagicMock() + session.kv_tasks = [] + + mid_slice = KVSlice( + is_last_slice=False, + block_ids_per_layer_groups=[np.array([0, 1], dtype=np.int64)], + ) + + transceiver = MagicMock() + transceiver._enable_pipelined_transfer = True + transceiver.kv_transfer_timeout_ms = None + transceiver._get_or_create_send_session.return_value = session + transceiver._build_prefill_chunk.return_value = mid_slice + + request = SimpleNamespace( + py_disaggregated_params=DisaggregatedParams(disagg_request_id=42), + request_id=42, + prompt_len=8, + py_beam_width=1, + py_kv_transfer_start_time=None, + set_kv_cache_transfer_start=lambda _ts: None, + ) + + KvCacheTransceiverV2.respond_and_send_async(transceiver, request) + + session.send.assert_called_once_with(mid_slice) + transceiver._finalize_send.assert_not_called() + + +# --------------------------------------------------------------------------- +# Retired send sessions +# --------------------------------------------------------------------------- + + +def _make_send_session_transceiver(sessions=None): + transceiver = MagicMock() + transceiver._send_sessions = dict(sessions or {}) + return transceiver + + +def _make_retirable_request(retired: bool, rid: int = 42): + return SimpleNamespace( + py_disaggregated_params=DisaggregatedParams(disagg_request_id=rid), + request_id=rid, + py_kv_send_session_retired=retired, + state=LlmRequestState.CONTEXT_INIT, + ) + + +def test_get_or_create_send_session_refuses_retired_request(): + """Closing a send session drops the peer registration, so a new one is inert. + + Its tasks would sit in INIT forever, which is neither completed nor failed, + so the request would never resolve and its blocks would stay pinned. + """ + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 + + transceiver = _make_send_session_transceiver() + request = _make_retirable_request(retired=True) + + session = KvCacheTransceiverV2._get_or_create_send_session(transceiver, request) + + assert session is None + assert request.state == LlmRequestState.DISAGG_TRANS_ERROR + transceiver._transfer_worker.create_tx_session.assert_not_called() + assert transceiver._send_sessions == {} + + +def test_get_or_create_send_session_creates_for_fresh_request(): + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 + + transceiver = _make_send_session_transceiver() + request = _make_retirable_request(retired=False) + created = transceiver._transfer_worker.create_tx_session.return_value + + session = KvCacheTransceiverV2._get_or_create_send_session(transceiver, request) + + assert session is created + assert transceiver._send_sessions == {42: created} + assert request.state == LlmRequestState.CONTEXT_INIT + + +def test_get_or_create_send_session_prefers_live_session_over_retired_flag(): + """A live session is the source of truth; the flag only bars re-creation.""" + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 + + live = MagicMock() + transceiver = _make_send_session_transceiver({42: live}) + request = _make_retirable_request(retired=True) + + session = KvCacheTransceiverV2._get_or_create_send_session(transceiver, request) + + assert session is live + assert request.state == LlmRequestState.CONTEXT_INIT + + +def test_respond_and_send_async_returns_early_when_session_refused(): + """A refused session must not build, send, or finalize anything.""" + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 + + transceiver = MagicMock() + transceiver._enable_pipelined_transfer = True + transceiver.kv_transfer_timeout_ms = None + transceiver._get_or_create_send_session.return_value = None + + request = SimpleNamespace( + py_disaggregated_params=DisaggregatedParams(disagg_request_id=42), + request_id=42, + prompt_len=8, + py_beam_width=1, + py_kv_transfer_start_time=None, + set_kv_cache_transfer_start=lambda _ts: None, + state=LlmRequestState.DISAGG_TRANS_ERROR, + ) + + KvCacheTransceiverV2.respond_and_send_async(transceiver, request) + + transceiver._build_prefill_chunk.assert_not_called() + transceiver._create_kv_slice.assert_not_called() + transceiver._finalize_send.assert_not_called() + assert request.state == LlmRequestState.DISAGG_TRANS_ERROR + + +# --------------------------------------------------------------------------- +# Context-side prefix reuse +# --------------------------------------------------------------------------- + +_REUSE_TPB = 4 +_REUSE_TOTAL_BLOCKS = 8 + + +def _build_prefill_chunk_for( + prepopulated_blocks, + chunk_start_block, + chunk_end_block, + resident_blocks=None, +): + """Drive the real _build_prefill_chunk for one chunk of a prefilling request. + + ``resident_blocks`` defaults to the chunk end, matching a source block list + that has only grown through the current chunk boundary. + """ + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 + + if resident_blocks is None: + resident_blocks = chunk_end_block + base_slice = KVSlice( + block_ids_per_layer_groups=[np.arange(resident_blocks, dtype=np.int64)], + ) + + transceiver = MagicMock() + transceiver._kv_cache_manager.tokens_per_block = _REUSE_TPB + transceiver._create_kv_slice.return_value = base_slice + transceiver._send_reqs = {} + + req = MagicMock() + req.py_disaggregated_params = DisaggregatedParams(disagg_request_id=42) + req.py_beam_width = 1 + req.prompt_len = _REUSE_TOTAL_BLOCKS * _REUSE_TPB + req.prepopulated_prompt_len = prepopulated_blocks * _REUSE_TPB + req.py_last_context_chunk = ( + chunk_start_block * _REUSE_TPB, + chunk_end_block * _REUSE_TPB, + ) + req.context_remaining_length = (_REUSE_TOTAL_BLOCKS - chunk_end_block) * _REUSE_TPB + + return KvCacheTransceiverV2._build_prefill_chunk(transceiver, req) + + +def test_first_chunk_covers_ctx_prefix_reuse(): + """The reused prefix is resident but no chunk spans it, so slice 0 extends to block 0.""" + kv_slice = _build_prefill_chunk_for( + prepopulated_blocks=3, + chunk_start_block=3, + chunk_end_block=6, + ) + + assert kv_slice.token_range == TokenRange(start=0, end=6 * _REUSE_TPB) + assert np.array_equal(kv_slice.block_ids_per_layer_groups[0], np.arange(6, dtype=np.int64)) + assert kv_slice.is_last_slice is False + + +@pytest.mark.parametrize( + "prepopulated_blocks,chunk_start_block,chunk_end_block,expected_start_block", + [ + (3, 6, 8, 6), + (0, 4, 8, 4), + (0, 0, 4, 0), + ], + ids=["after_reuse_hit", "no_reuse_later_chunk", "no_reuse_first_chunk"], +) +def test_only_the_first_chunk_extends_to_block_zero( + prepopulated_blocks, chunk_start_block, chunk_end_block, expected_start_block +): + """Chunks past the first keep their own start; without reuse nothing changes.""" + kv_slice = _build_prefill_chunk_for( + prepopulated_blocks=prepopulated_blocks, + chunk_start_block=chunk_start_block, + chunk_end_block=chunk_end_block, + resident_blocks=_REUSE_TOTAL_BLOCKS, + ) + + assert kv_slice.token_range == TokenRange( + start=expected_start_block * _REUSE_TPB, end=chunk_end_block * _REUSE_TPB + ) + assert np.array_equal( + kv_slice.block_ids_per_layer_groups[0], + np.arange(expected_start_block, chunk_end_block, dtype=np.int64), + ) + + +def test_single_chunk_with_reuse_degenerates_to_monolithic_slice(): + """One chunk plus a reuse hit yields the same slice shape a monolithic send would. + + token_range.start == 0 with is_last_slice makes _build_kv_write_meta take its + non-chunked branch, so the write is addressed exactly as an unpipelined one. + """ + kv_slice = _build_prefill_chunk_for( + prepopulated_blocks=3, + chunk_start_block=3, + chunk_end_block=_REUSE_TOTAL_BLOCKS, + resident_blocks=_REUSE_TOTAL_BLOCKS, + ) + + assert kv_slice.is_last_slice is True + assert kv_slice.token_range == TokenRange(start=0, end=_REUSE_TOTAL_BLOCKS * _REUSE_TPB) + assert np.array_equal( + kv_slice.block_ids_per_layer_groups[0], + np.arange(_REUSE_TOTAL_BLOCKS, dtype=np.int64), + ) + + +# --------------------------------------------------------------------------- +# Transfer activity as a dimension owned by the transceiver +# --------------------------------------------------------------------------- + + +def _make_transfer_state_transceiver(session=None, rid: int = 42): + """Transceiver stub whose session maps are the ownership record.""" + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 + + sessions = {rid: session} if session is not None else {} + transceiver = SimpleNamespace( + _wait_reqs={}, + _send_sessions=dict(sessions), + _send_reqs={rid: MagicMock()} if session is not None else {}, + _recv_sessions={}, + _recv_reqs={}, + ) + # Real teardown, so the predicate is checked against the actual bookkeeping. + transceiver._retire_send_session = MethodType( + KvCacheTransceiverV2._retire_send_session, transceiver + ) + return transceiver + + +def _make_transfer_state_request(rid=42, request_id: int = 42): + return SimpleNamespace( + py_disaggregated_params=( + DisaggregatedParams(disagg_request_id=rid) if rid is not None else None + ), + request_id=request_id, + ) + + +def test_has_inflight_transfer_tracks_send_session_lifetime(): + """Session membership answers the predicate, before and after teardown.""" + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 + + session = MagicMock() + session.has_transferring_tasks.return_value = False + transceiver = _make_transfer_state_transceiver(session) + request = _make_transfer_state_request() + + assert KvCacheTransceiverV2.has_inflight_transfer(transceiver, request) + assert KvCacheTransceiverV2.has_any_inflight_transfer(transceiver) + + assert KvCacheTransceiverV2.cancel_request(transceiver, request) + + assert not KvCacheTransceiverV2.has_inflight_transfer(transceiver, request) + assert not KvCacheTransceiverV2.has_any_inflight_transfer(transceiver) + + +def test_has_inflight_transfer_survives_mid_write_cancel(): + """A cancel that cannot complete keeps the ownership record alive.""" + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 + + session = MagicMock() + session.has_transferring_tasks.return_value = True + transceiver = _make_transfer_state_transceiver(session) + request = _make_transfer_state_request() + + assert not KvCacheTransceiverV2.cancel_request(transceiver, request) + assert KvCacheTransceiverV2.has_inflight_transfer(transceiver, request) + + +def test_has_inflight_transfer_false_without_disagg_params(): + """A request that never registered a session owns no transfer resources.""" + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 + + transceiver = _make_transfer_state_transceiver() + request = _make_transfer_state_request(rid=None, request_id=7) + + assert not KvCacheTransceiverV2.has_inflight_transfer(transceiver, request) + + +def test_is_request_in_transmission_uses_transceiver_predicate(): + """A mid-prefill request still counts as transmitting despite CONTEXT_INIT.""" + from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor + + executor = MagicMock() + executor.kv_cache_transceiver.has_inflight_transfer.return_value = True + + request = SimpleNamespace(state=LlmRequestState.CONTEXT_INIT) + + assert PyExecutor._is_request_in_transmission(executor, request) + executor.kv_cache_transceiver.has_inflight_transfer.assert_called_once_with(request) + + +def test_is_request_in_transmission_false_when_nothing_in_flight(): + from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor + + executor = MagicMock() + executor.kv_cache_transceiver.has_inflight_transfer.return_value = False + + request = SimpleNamespace(state=LlmRequestState.CONTEXT_INIT) + + assert not PyExecutor._is_request_in_transmission(executor, request) + + +def test_try_cancel_request_propagates_mid_write_failure(): + """Cancelling mid-prefill delegates and reports the retry-needed result.""" + from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor + + executor = MagicMock() + executor._is_request_in_transmission.return_value = True + executor._is_disagg_inflight_cancel_active.return_value = False + executor.kv_cache_transceiver.cancel_request.return_value = False + # The delegation goes through _request_kv_transfer_cancellation, so run the + # real helper to keep the assertion on the transceiver call itself. + executor._request_kv_transfer_cancellation = ( + lambda req: PyExecutor._request_kv_transfer_cancellation(executor, req) + ) + + request = SimpleNamespace(state=LlmRequestState.CONTEXT_INIT) + + assert not PyExecutor._try_cancel_request(executor, request) + executor.kv_cache_transceiver.cancel_request.assert_called_once_with(request) + + +def _make_send_kv_executor(canceled_req_ids, retired: bool = False): + executor = MagicMock() + executor.kv_connector_manager = None + executor.canceled_req_ids = list(canceled_req_ids) + executor.kv_cache_transceiver.pipeline_transfer_enabled = True + executor.kv_cache_transceiver.kv_transfer_timeout_ms = None + executor.kv_cache_transceiver.has_retired_send_session.return_value = retired + return executor + + +def _make_send_kv_request( + is_last_chunk: bool, + request_id: int = 7, + state=LlmRequestState.CONTEXT_INIT, +): + return SimpleNamespace( + is_context_only_request=True, + is_finished_due_to_cancellation=False, + is_context_finished=is_last_chunk, + is_finished_due_to_length=False, + is_child=False, + parent_request_id=None, + py_request_id=request_id, + py_kv_transfer_start_time=None, + state=state, + ) + + +def test_send_kv_async_skips_intermediate_chunk_for_cancelled_request(): + """A cancelled session must not be fed another chunk.""" + from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor + + executor = _make_send_kv_executor([7]) + request = _make_send_kv_request(is_last_chunk=False) + + PyExecutor._send_kv_async(executor, [request]) + + executor.kv_cache_transceiver.respond_and_send_async.assert_not_called() + + +def test_send_kv_async_sends_intermediate_chunk_when_not_cancelled(): + from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor + + executor = _make_send_kv_executor([]) + request = _make_send_kv_request(is_last_chunk=False) + + PyExecutor._send_kv_async(executor, [request]) + + executor.kv_cache_transceiver.respond_and_send_async.assert_called_once_with(request) + + +def test_send_kv_async_skips_intermediate_chunk_for_failed_request(): + """An error path already failed and freed this request, leaving its chunk bounds unset. + + _update_request_states skips GENERATION_COMPLETE requests, so + py_last_context_chunk is still (None, None) and building a chunk from it + would fault. The request stays in scheduled_requests either way. + """ + from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor + + executor = _make_send_kv_executor([]) + request = _make_send_kv_request(is_last_chunk=False, state=LlmRequestState.GENERATION_COMPLETE) + + PyExecutor._send_kv_async(executor, [request]) + + executor.kv_cache_transceiver.respond_and_send_async.assert_not_called() + + +def test_send_kv_async_skips_retired_request_mid_prefill(): + """A retired session cannot reach its peer, so the request is failed instead.""" + from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor + + executor = _make_send_kv_executor([], retired=True) + request = _make_send_kv_request(is_last_chunk=False) + + PyExecutor._send_kv_async(executor, [request]) + + assert request.state == LlmRequestState.DISAGG_TRANS_ERROR + executor.kv_cache_transceiver.respond_and_send_async.assert_not_called() + + +def test_send_kv_async_skips_retired_request_before_start_transfer(): + """The gate precedes start_transfer, which pins blocks only end_transfer releases. + + Reaching the final-chunk branch would register the request with the transfer + manager for a transfer that can never complete. + """ + from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor + + executor = _make_send_kv_executor([], retired=True) + request = _make_send_kv_request(is_last_chunk=True) + + PyExecutor._send_kv_async(executor, [request]) + + assert request.state == LlmRequestState.DISAGG_TRANS_ERROR + executor.async_transfer_manager.start_transfer.assert_not_called() + executor.kv_cache_transceiver.respond_and_send_async.assert_not_called() + + +def test_send_kv_async_still_sends_final_chunk_for_cancelled_request(): + """The final chunk stays unconditional so nothing strands in the manager.""" + from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor + + executor = _make_send_kv_executor([7]) + request = _make_send_kv_request(is_last_chunk=True) + + PyExecutor._send_kv_async(executor, [request]) + + executor.async_transfer_manager.start_transfer.assert_called_once_with(request) + executor.kv_cache_transceiver.respond_and_send_async.assert_called_once_with(request) + + +def _make_timeout_request(request_id: int = 7, elapsed_s: float = 10.0): + return SimpleNamespace( + is_context_only_request=True, + is_disagg_generation_transmission_in_progress=False, + py_request_id=request_id, + py_kv_transfer_start_time=time.monotonic() - elapsed_s, + py_kv_transfer_timed_out=False, + state=LlmRequestState.CONTEXT_INIT, + ) + + +def test_check_kv_transfer_timeout_flags_context_request_in_transfer(): + """Context requests are monitored via the transfer manager. + + They enter it on the last chunk, at the same time as their clock starts. + """ + from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor + + executor = MagicMock() + executor.kv_cache_transceiver.kv_transfer_timeout_ms = 100 + request = _make_timeout_request() + executor.async_transfer_manager.requests_in_transfer.return_value = {7: request} + executor.active_requests = [request] + + PyExecutor._check_kv_transfer_timeout(executor) + + assert request.py_kv_transfer_timed_out + + +def test_check_kv_transfer_timeout_ignores_context_request_not_in_transfer(): + """A request still being prefilled is not monitored. + + It holds its KV pages because it is computing, not because a chunk is in + flight, so the timer only needs to cover the final transfer. + """ + from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor + + executor = MagicMock() + executor.kv_cache_transceiver.kv_transfer_timeout_ms = 100 + executor.async_transfer_manager.requests_in_transfer.return_value = {} + + request = _make_timeout_request() + executor.active_requests = [request] + + PyExecutor._check_kv_transfer_timeout(executor) + + assert not request.py_kv_transfer_timed_out + + +def test_has_any_inflight_kv_transfer_ors_in_transceiver(): + from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor + + executor = MagicMock() + executor.async_transfer_manager.has_any_inflight_requests.return_value = False + executor.kv_cache_transceiver.has_any_inflight_transfer.return_value = True + + assert PyExecutor._has_any_inflight_kv_transfer(executor) + + executor.kv_cache_transceiver.has_any_inflight_transfer.return_value = False + assert not PyExecutor._has_any_inflight_kv_transfer(executor) + + +def test_ctx_transfer_status_leaves_mid_prefill_request_alone(): + """The timeout path never cancels a request still being prefilled. + + py_kv_transfer_timed_out can only be set for requests the transfer manager + already knows about, so the mid-prefill case has nothing to act on. + """ + from tensorrt_llm._torch.pyexecutor.py_executor import PyExecutor + + executor = MagicMock() + executor.kv_cache_transceiver.check_context_transfer_status.return_value = ([], []) + executor.async_transfer_manager.requests_in_transfer.return_value = {} + + request = _make_timeout_request() + request.py_kv_transfer_timed_out = True + executor.active_requests = [request] + + PyExecutor._check_disagg_ctx_cache_transfer_status(executor) + + executor.kv_cache_transceiver.cancel_request.assert_not_called() + assert request.state == LlmRequestState.CONTEXT_INIT diff --git a/tests/unittest/disaggregated/test_disagg_utils.py b/tests/unittest/disaggregated/test_disagg_utils.py index 42a0cce50de7..a3d8f2d0ec9a 100644 --- a/tests/unittest/disaggregated/test_disagg_utils.py +++ b/tests/unittest/disaggregated/test_disagg_utils.py @@ -110,6 +110,35 @@ def test_parse_disagg_config_file(sample_yaml_file, sample_yaml_config): verify_disagg_config(config, sample_yaml_config) +def test_parse_disagg_config_file_rejects_empty_file(tmp_path): + config_file = tmp_path / "empty.yaml" + config_file.write_text("") + + with pytest.raises(ValueError, match="Disaggregated config file is empty"): + parse_disagg_config_file(config_file) + + +def test_parse_disagg_config_file_validates_schedule_style_override(tmp_path): + config_file = tmp_path / "pipelined.yaml" + config_file.write_text( + yaml.safe_dump({ + "schedule_style": "generation_first", + "context_servers": { + "cache_transceiver_config": { + "enable_pipelined_transfer": True, + }, + }, + "generation_servers": {}, + })) + + with pytest.raises( + ValueError, + match="enable_pipelined_transfer=True requires top-level " + "schedule_style='generation_first'"): + parse_disagg_config_file(config_file, + schedule_style_override="context_first") + + @pytest.mark.parametrize("sample_yaml_config", ["disagg_cluster", ""], indirect=True) def test_extract_disagg_cfg(sample_yaml_config): @@ -234,6 +263,70 @@ def test_extract_disagg_cfg_rejects_non_string_internal_request_auth_key(): extract_disagg_cfg(internal_request_auth_key=123) +@pytest.mark.parametrize("server_group", + ["context_servers", "generation_servers"]) +def test_extract_disagg_cfg_rejects_pipelined_transfer_with_context_first( + server_group): + server_configs = { + "context_servers": {}, + "generation_servers": {}, + } + server_configs[server_group] = { + "cache_transceiver_config": { + "enable_pipelined_transfer": True, + }, + } + + with pytest.raises( + ValueError, + match="enable_pipelined_transfer=True requires top-level " + "schedule_style='generation_first'"): + extract_disagg_cfg(schedule_style="context_first", **server_configs) + + +def test_extract_disagg_cfg_allows_pipelined_transfer_with_generation_first(): + config = extract_disagg_cfg( + schedule_style="generation_first", + context_servers={ + "cache_transceiver_config": { + "enable_pipelined_transfer": True, + }, + }, + generation_servers={}, + ) + + assert config.schedule_style == "generation_first" + + +def test_extract_disagg_cfg_rejects_invalid_schedule_style(): + with pytest.raises(ValueError, match="schedule_style must be one of"): + extract_disagg_cfg( + schedule_style="generation-first", + context_servers={}, + generation_servers={}, + ) + + +@pytest.mark.parametrize( + "cache_transceiver_config, expected_error", + [ + ({ + "enable_pipelined_transfer": 1 + }, "enable_pipelined_transfer must be a boolean"), + ([], "cache_transceiver_config must be a mapping"), + ], +) +def test_extract_disagg_cfg_rejects_invalid_cache_transceiver_config( + cache_transceiver_config, expected_error): + with pytest.raises(ValueError, match=expected_error): + extract_disagg_cfg( + context_servers={ + "cache_transceiver_config": cache_transceiver_config, + }, + generation_servers={}, + ) + + def test_extract_ctx_gen_cfgs(): configs = extract_ctx_gen_cfgs( type="ctx", diff --git a/tests/unittest/disaggregated/test_kv_transfer.py b/tests/unittest/disaggregated/test_kv_transfer.py index 8abd9147dc72..6039fbf1b8da 100644 --- a/tests/unittest/disaggregated/test_kv_transfer.py +++ b/tests/unittest/disaggregated/test_kv_transfer.py @@ -4,6 +4,8 @@ import random import time import uuid +from types import SimpleNamespace +from unittest.mock import MagicMock # Force a deterministic UCX config regardless of what the cluster/CI injects # (the CI agent bootstrap exports UCX_TLS=tcp,cuda_copy,cuda_ipc before pytest @@ -46,6 +48,7 @@ ) from tensorrt_llm._torch.disaggregation.native.transfer import TransferWorker, TransferWorkerConfig from tensorrt_llm._torch.disaggregation.resource.kv_extractor import KVRegionExtractorV1 +from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2 from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, LlmRequestType from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager @@ -181,6 +184,242 @@ def test_session_status_enum(): assert len(SessionStatus) == 7 +# --------------------------------------------------------------------------- +# Pipelined prefill chunk creation tests +# --------------------------------------------------------------------------- + + +def _send_prefill_chunks( + all_block_ids, + chunk_size_blocks, + mamba_state_index=None, + tokens_per_block=1, + sender_session=None, + prepopulated_blocks=0, +): + """Build (and optionally send) chunks through the real ``_build_prefill_chunk`` path. + + ``_build_prefill_chunk`` derives the chunk bounds from ``req.py_last_context_chunk`` + and ``req.context_remaining_length`` and returns a ``KVSlice`` (``respond_and_send_async`` + is responsible for sending). This helper + drives one call per chunk, collecting the returned slices; when ``sender_session`` is + provided it forwards each built slice to the real session to mirror the send path. + + ``prepopulated_blocks`` models a context-side prefix-reuse hit: the scheduler's + chunks start after the reused prefix, while the source block list still holds + every block. The first slice must cover the reused prefix anyway. + """ + all_block_ids = [np.asarray(ids, dtype=np.int64) for ids in all_block_ids] + total_blocks = max((len(ids) for ids in all_block_ids), default=0) + base_slice = KVSlice( + block_ids_per_layer_groups=all_block_ids, + mamba_state_index=mamba_state_index, + ) + session = sender_session if sender_session is not None else MagicMock() + session.kv_tasks = [] + transceiver = MagicMock() + transceiver._get_or_create_send_session.return_value = session + transceiver._create_kv_slice = MagicMock(return_value=base_slice) + transceiver._reuse_adapter.tokens_per_block = tokens_per_block + transceiver._kv_cache_manager.tokens_per_block = tokens_per_block + transceiver._kv_cache_manager.kv_cache_map = {} + transceiver._send_reqs = {} + + prompt_len = total_blocks * tokens_per_block + req = MagicMock() + req.py_disaggregated_params = DisaggregatedParams(disagg_request_id=42) + req.prompt_len = prompt_len + req.py_beam_width = 1 + # Explicit: a MagicMock attribute never compares equal to the chunk start, so + # _build_prefill_chunk's first-chunk branch would never be exercised. + req.prepopulated_prompt_len = prepopulated_blocks * tokens_per_block + + first_block = min(prepopulated_blocks, total_blocks) + if chunk_size_blocks is None or chunk_size_blocks >= total_blocks - first_block: + chunk_ranges = [(first_block, total_blocks)] + else: + chunk_ranges = [ + (start, min(start + chunk_size_blocks, total_blocks)) + for start in range(first_block, total_blocks, chunk_size_blocks) + ] + if not chunk_ranges: + chunk_ranges = [(0, 0)] + + slices = [] + for idx, (chunk_start, chunk_end) in enumerate(chunk_ranges): + is_last_chunk = idx == len(chunk_ranges) - 1 + req.py_last_context_chunk = ( + chunk_start * tokens_per_block, + chunk_end * tokens_per_block, + ) + req.context_remaining_length = ( + 0 if is_last_chunk else (total_blocks - chunk_end) * tokens_per_block + ) + kv_slice = KvCacheTransceiverV2._build_prefill_chunk(transceiver, req) + slices.append(kv_slice) + if sender_session is not None: + session.send(kv_slice) + + if sender_session is not None: + return [] + return slices + + +def test_build_prefill_chunk_projects_incremental_source_against_full_prompt(): + """A growing source uses the current chunk end while retaining the full prompt span.""" + tokens_per_block = 128 + prompt_blocks = 1000 + chunk_blocks = 16 + + transceiver = MagicMock() + transceiver._reuse_adapter.tokens_per_block = tokens_per_block + transceiver._kv_cache_manager.tokens_per_block = tokens_per_block + transceiver._kv_cache_manager.kv_cache_map = {} + transceiver._send_reqs = {} + + req = MagicMock() + req.py_disaggregated_params = DisaggregatedParams(disagg_request_id=42) + req.py_request_id = 42 + req.prompt_len = prompt_blocks * tokens_per_block + req.py_beam_width = 1 + req.prepopulated_prompt_len = 0 + + for chunk_idx in range(2): + chunk_start = chunk_idx * chunk_blocks + chunk_end = chunk_start + chunk_blocks + resident_ids = np.arange(chunk_end, dtype=np.int64) + transceiver._create_kv_slice.return_value = KVSlice( + block_ids_per_layer_groups=[resident_ids] + ) + req.py_last_context_chunk = ( + chunk_start * tokens_per_block, + chunk_end * tokens_per_block, + ) + req.context_remaining_length = req.prompt_len - req.py_last_context_chunk[1] + + kv_slice = KvCacheTransceiverV2._build_prefill_chunk(transceiver, req) + + assert np.array_equal( + kv_slice.block_ids_per_layer_groups[0], + np.arange(chunk_start, chunk_end, dtype=np.int64), + ) + assert kv_slice.total_blocks == prompt_blocks + + +@pytest.mark.parametrize( + "source_block_ids", + [ + np.arange(16, dtype=np.int64), + np.arange(9, 13, dtype=np.int64), + ], + ids=["v1_full_prompt_allocation", "v2_incremental_allocation"], +) +def test_build_prefill_chunk_normalizes_swa_source_to_computed_prefix(source_block_ids): + """A partial SWA chunk must not select full-prompt pages beyond its computed end.""" + tokens_per_block = 8 + prompt_blocks = 16 + window_blocks = 4 + + layer_group = SimpleNamespace(sliding_window_size=window_blocks * tokens_per_block) + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._reuse_adapter = SimpleNamespace( + tokens_per_block=tokens_per_block, + get_cached_token_count_per_layer_group=lambda req, layer_groups: [0], + get_block_ids=lambda req, idx, lg: source_block_ids, + ) + transceiver._page_table = SimpleNamespace(layer_groups=[layer_group]) + transceiver._kv_cache_manager = SimpleNamespace( + tokens_per_block=tokens_per_block, + num_extra_kv_tokens=0, + ) + transceiver._send_reqs = {} + + req = MagicMock() + req.py_disaggregated_params = DisaggregatedParams(disagg_request_id=42) + req.py_request_id = 42 + req.prompt_len = prompt_blocks * tokens_per_block + req.py_beam_width = 1 + req.prepopulated_prompt_len = 0 + req.py_last_context_chunk = (11 * tokens_per_block, 13 * tokens_per_block) + req.context_remaining_length = 3 * tokens_per_block + req.is_generation_only_request.return_value = False + + kv_slice = KvCacheTransceiverV2._build_prefill_chunk(transceiver, req) + + np.testing.assert_array_equal(kv_slice.block_ids_per_layer_groups[0], [12]) + + +@pytest.mark.parametrize( + "all_block_ids,chunk_size_blocks,expected_num_slices", + [ + ([[0, 1, 2, 3, 4, 5, 6, 7]], None, 1), + ([[0, 1, 2, 3, 4, 5, 6, 7]], 4, 2), + ([list(range(10))], 4, 3), + ([[], []], 4, 1), + ([[0, 1, 2]], 64, 1), + ], + ids=["no_chunking", "even_split", "uneven_split", "empty_blocks", "chunk_larger_than_total"], +) +def test_send_prefill_chunks_basic(all_block_ids, chunk_size_blocks, expected_num_slices): + """Pipelined prefill chunking produces the expected number of slices.""" + slices = _send_prefill_chunks(all_block_ids, chunk_size_blocks) + assert len(slices) == expected_num_slices + assert slices[-1].is_last_slice is True + if expected_num_slices > 1: + for s in slices[:-1]: + assert s.is_last_slice is False + + +@pytest.mark.parametrize("prepopulated_blocks", [0, 1, 5], ids=["no_reuse", "reuse_1", "reuse_5"]) +def test_send_prefill_chunks_integrity_check(prepopulated_blocks): + """Every block reaches exactly one slice, including a reused context-side prefix. + + With a reuse hit the scheduler's chunks start past the prefix, so this is the + check that catches blocks that no chunk covers. + """ + all_block_ids = [list(range(17)), list(range(17))] + slices = _send_prefill_chunks( + all_block_ids, + chunk_size_blocks=4, + prepopulated_blocks=prepopulated_blocks, + ) + for lg_idx, original in enumerate(all_block_ids): + reassembled = [] + for s in slices: + reassembled.extend(s.block_ids_per_layer_groups[lg_idx]) + assert reassembled == original + + +def test_send_prefill_chunks_multiple_layer_groups(): + """Each source layer group is a resident suffix ending at the current chunk.""" + all_block_ids = [list(range(8)), list(range(3))] + slices = _send_prefill_chunks(all_block_ids, chunk_size_blocks=4) + assert len(slices) == 2 + assert np.array_equal(slices[0].block_ids_per_layer_groups[0], np.array([0, 1, 2, 3])) + assert np.array_equal(slices[1].block_ids_per_layer_groups[0], np.array([4, 5, 6, 7])) + assert np.array_equal(slices[0].block_ids_per_layer_groups[1], np.array([0, 1, 2])) + assert np.array_equal(slices[1].block_ids_per_layer_groups[1], np.array([0, 1, 2])) + assert slices[0].token_range == TokenRange(start=0, end=4) + assert slices[1].token_range == TokenRange(start=4, end=8) + + +def test_send_prefill_chunks_preserves_mamba_state_index(): + """mamba_state_index is propagated to every chunk slice.""" + all_block_ids = [list(range(8))] + slices = _send_prefill_chunks(all_block_ids, chunk_size_blocks=4, mamba_state_index=42) + assert len(slices) == 2 + for s in slices: + assert s.mamba_state_index == 42 + + +def test_send_prefill_chunks_none_mamba_state_index(): + """mamba_state_index=None is preserved when not set.""" + all_block_ids = [list(range(4))] + slices = _send_prefill_chunks(all_block_ids, chunk_size_blocks=4) + assert len(slices) == 1 + assert slices[0].mamba_state_index is None + + def create_transfer_worker_setup( ctx_tp: int, ctx_pp: int, @@ -1115,7 +1354,12 @@ def test_transfer_worker_v2_with_window( @pytest.mark.timeout(120) @pytest.mark.parametrize("use_v2", [False, True], ids=["v1", "v2"]) -def test_transfer_with_gen_prefix_offset(use_v2): +@pytest.mark.parametrize( + "chunk_size_blocks", + [None, 2], + ids=["single_slice", "sender_chunked"], +) +def test_transfer_with_gen_prefix_offset(use_v2, chunk_size_blocks): """Verify that only suffix blocks are transferred when gen has a prefix offset. Simulates gen-side prefix cache: ctx sends all blocks for [0, request_len), @@ -1204,13 +1448,7 @@ def test_transfer_with_gen_prefix_offset(use_v2): ] try: - # Ctx sends all blocks tx = ctx_tw.create_tx_session(ctx_request) - send_slice = KVSlice( - is_last_slice=True, - block_ids_per_layer_groups=ctx_block_ids, - token_range=TokenRange(start=0, end=request_len), - ) # Gen receives only the suffix list; dst_start is derived from block count. rx = gen_tw.create_rx_session(gen_request) @@ -1220,7 +1458,22 @@ def test_transfer_with_gen_prefix_offset(use_v2): token_range=TokenRange(start=0, end=request_len), ) rx.receive(recv_slice) - tx.send(send_slice) + + if chunk_size_blocks is None: + tx.send( + KVSlice( + is_last_slice=True, + block_ids_per_layer_groups=ctx_block_ids, + token_range=TokenRange(start=0, end=request_len), + ) + ) + else: + _send_prefill_chunks( + ctx_block_ids, + chunk_size_blocks=chunk_size_blocks, + tokens_per_block=tokens_per_block, + sender_session=tx, + ) result = tx.wait_complete() assert result == WaitResult.COMPLETED, f"tx wait_complete returned {result}" @@ -1318,7 +1571,7 @@ def test_session_cancel_before_send(): @pytest.mark.timeout(60) def test_session_cancel_after_send(): - """TxSession cancelled after send() queues INIT tasks; future raises.""" + """TxSession cancelled after send() queues INIT tasks fails the event wait.""" tensorrt_llm.logger.set_level("debug") setup = create_transfer_worker_setup( ctx_tp=1, @@ -1353,19 +1606,229 @@ def test_session_cancel_after_send(): page_table = ctx_transfer_worker._rank_info.page_table block_ids_per_groups = [np.array([], dtype=np.int64) for _ in page_table.layer_groups] kv_slice = KVSlice(is_last_slice=True, block_ids_per_layer_groups=block_ids_per_groups) - future = tx_session.send(kv_slice) + tx_session.send(kv_slice) # No receiver registered yet; task is INIT. tx_session.cancel() assert tx_session.status == SessionStatus.CANCELLED assert tx_session.has_failed() - # Future for the cancelled INIT task must raise. - with pytest.raises(Exception): - future.result(timeout=5.0) + assert tx_session.wait_complete() == WaitResult.FAILED tx_session.close() finally: ctx_transfer_worker.shutdown() + + +def _setup_chunked_request(setup, ctx_request_id, gen_request_id, request_len): + """Create requests, allocate KV, and collect block IDs for chunked transfer tests.""" + ctx_transfer_workers = setup["ctx_transfer_workers"] + ctx_kv_cache_managers = setup["ctx_kv_cache_managers"] + gen_transfer_workers = setup["gen_transfer_workers"] + gen_kv_cache_managers = setup["gen_kv_cache_managers"] + ctx_info_endpoint = setup["ctx_info_endpoint"] + use_v2 = setup["use_v2"] + tokens_per_block = setup["tokens_per_block"] + + sampling_params = SamplingParams() + unique_rid = uuid.uuid4().int & 0x7FFFFFFFFFFFFFFF + + ctx_request = LlmRequest( + request_id=ctx_request_id, + max_new_tokens=1, + input_tokens=list(range(request_len)), + sampling_config=tensorrt_llm.bindings.SamplingConfig( + sampling_params._get_sampling_config() + ), + is_streaming=False, + llm_request_type=LlmRequestType.LLMREQUEST_TYPE_CONTEXT_ONLY, + ) + ctx_request.py_disaggregated_params = DisaggregatedParams(disagg_request_id=unique_rid) + + gen_request = LlmRequest( + request_id=gen_request_id, + max_new_tokens=1, + input_tokens=list(range(request_len)), + sampling_config=tensorrt_llm.bindings.SamplingConfig( + sampling_params._get_sampling_config() + ), + is_streaming=False, + llm_request_type=LlmRequestType.LLMREQUEST_TYPE_GENERATION_ONLY, + ) + gen_request.py_disaggregated_params = DisaggregatedParams( + ctx_request_id=ctx_request.py_request_id, + ctx_dp_rank=0, + ctx_info_endpoint=ctx_info_endpoint, + disagg_request_id=unique_rid, + ) + + ctx_kv_caches, gen_kv_caches = [], [] + for mgr in ctx_kv_cache_managers: + if use_v2: + kv = mgr._create_kv_cache(ctx_request.py_request_id, None, None) + assert kv.resume(torch.cuda.current_stream().cuda_stream) + assert kv.resize(request_len) + ctx_kv_caches.append(kv) + else: + mgr.impl.add_sequence_batch( + [(ctx_request.py_request_id, request_len, 1)], [ctx_request] + ) + + for mgr in gen_kv_cache_managers: + if use_v2: + kv = mgr._create_kv_cache(gen_request.py_request_id, None, None) + assert kv.resume(torch.cuda.current_stream().cuda_stream) + assert kv.resize(request_len) + gen_kv_caches.append(kv) + else: + mgr.impl.add_sequence_batch( + [(gen_request.py_request_id, request_len, 1)], [gen_request] + ) + + ctx_block_ids = [ + get_block_ids_per_layer_groups(mgr, tw, ctx_request.py_request_id, use_v2, tokens_per_block) + for mgr, tw in zip(ctx_kv_cache_managers, ctx_transfer_workers, strict=True) + ] + gen_block_ids = [ + get_block_ids_per_layer_groups(mgr, tw, gen_request.py_request_id, use_v2, tokens_per_block) + for mgr, tw in zip(gen_kv_cache_managers, gen_transfer_workers, strict=True) + ] + + return { + "ctx_request": ctx_request, + "gen_request": gen_request, + "ctx_kv_caches": ctx_kv_caches, + "gen_kv_caches": gen_kv_caches, + "ctx_block_ids": ctx_block_ids, + "gen_block_ids": gen_block_ids, + } + + +def _verify_and_cleanup_chunked(setup, ctx_info, sender_sessions, receiver_sessions): + """Shared verification and cleanup for chunked transfer tests.""" + ctx_kv_cache_managers = setup["ctx_kv_cache_managers"] + gen_kv_cache_managers = setup["gen_kv_cache_managers"] + use_v2 = setup["use_v2"] + + ctx_block_ids = ctx_info["ctx_block_ids"] + gen_block_ids = ctx_info["gen_block_ids"] + + for session in sender_sessions: + assert session.status == SessionStatus.KV_TRANSFERRED + for session in receiver_sessions: + assert session.status == SessionStatus.KV_TRANSFERRED + + num_layer_groups = len(ctx_block_ids[0]) + for lg_id in range(num_layer_groups): + ctx_data = [ + get_block_data(mgr, bids[lg_id], lg_id, use_v2, ctx_info["ctx_request"].py_request_id) + for mgr, bids in zip(ctx_kv_cache_managers, ctx_block_ids, strict=True) + ] + gen_data = [ + get_block_data(mgr, bids[lg_id], lg_id, use_v2, ctx_info["gen_request"].py_request_id) + for mgr, bids in zip(gen_kv_cache_managers, gen_block_ids, strict=True) + ] + for c, g in zip(ctx_data, gen_data, strict=True): + assert c.equal(g), f"Layer group {lg_id}: data mismatch with chunked transfer" + + for s in receiver_sessions: + s.close() + for s in sender_sessions: + s.close() + if use_v2: + torch.cuda.current_stream().synchronize() + for kv in ctx_info["ctx_kv_caches"]: + kv.close() + for kv in ctx_info["gen_kv_caches"]: + kv.close() + + +def add_and_verify_chunked_request( + setup, + ctx_request_id, + gen_request_id, + request_len, + chunk_size_blocks, +): + """Chunked transfer variant: sender sends N slices, receiver sends 1.""" + ctx_transfer_workers = setup["ctx_transfer_workers"] + gen_transfer_workers = setup["gen_transfer_workers"] + + ctx_info = _setup_chunked_request(setup, ctx_request_id, gen_request_id, request_len) + ctx_block_ids = ctx_info["ctx_block_ids"] + gen_block_ids = ctx_info["gen_block_ids"] + token_range = TokenRange(start=0, end=request_len) + + sender_sessions = [tw.create_tx_session(ctx_info["ctx_request"]) for tw in ctx_transfer_workers] + for sender_session, block_ids_per_groups in zip(sender_sessions, ctx_block_ids, strict=True): + _send_prefill_chunks( + block_ids_per_groups, + chunk_size_blocks=chunk_size_blocks, + tokens_per_block=setup["tokens_per_block"], + sender_session=sender_session, + ) + + receiver_sessions = [ + tw.create_rx_session(ctx_info["gen_request"]) for tw in gen_transfer_workers + ] + for recv_session, block_ids_per_groups in zip(receiver_sessions, gen_block_ids, strict=True): + full_slice = KVSlice( + is_last_slice=True, + block_ids_per_layer_groups=block_ids_per_groups, + token_range=token_range, + ) + recv_session.receive(full_slice) + + for session in sender_sessions: + result = session.wait_complete() + assert result == WaitResult.COMPLETED, f"tx wait_complete returned {result}" + for session in receiver_sessions: + result = session.wait_complete(blocking=True) + assert result == WaitResult.COMPLETED, f"rx wait_complete returned {result}" + + _verify_and_cleanup_chunked(setup, ctx_info, sender_sessions, receiver_sessions) + + +CHUNKED_TEST_CONFIGS = [ + (1, 1, False, 1, 1, False, False, True, "v2_tp1_pp1_chunked"), + (1, 1, False, 1, 1, False, False, False, "v1_tp1_pp1_chunked"), +] + + +@pytest.mark.timeout(120) +@pytest.mark.parametrize( + "ctx_tp,ctx_pp,ctx_enable_dp,gen_tp,gen_pp,gen_enable_dp,is_mla,use_v2", + [(c[0], c[1], c[2], c[3], c[4], c[5], c[6], c[7]) for c in CHUNKED_TEST_CONFIGS], + ids=[c[8] for c in CHUNKED_TEST_CONFIGS], +) +def test_transfer_worker_chunked( + ctx_tp, ctx_pp, ctx_enable_dp, gen_tp, gen_pp, gen_enable_dp, is_mla, use_v2 +): + """Test transfer worker with sender-side chunking for V1 and V2.""" + tensorrt_llm.logger.set_level("info") + logger.info(f"Test transfer worker {'V2' if use_v2 else 'V1'} with chunked transfer") + + setup = create_transfer_worker_setup( + ctx_tp=ctx_tp, + ctx_pp=ctx_pp, + ctx_enable_dp=ctx_enable_dp, + gen_tp=gen_tp, + gen_pp=gen_pp, + gen_enable_dp=gen_enable_dp, + is_mla=is_mla, + use_v2=use_v2, + ) + + request_len = setup["request_len"] + tokens_per_block = setup["tokens_per_block"] + total_blocks = (request_len + tokens_per_block - 1) // tokens_per_block + chunk_size = max(1, total_blocks // 2) + + try: + add_and_verify_chunked_request(setup, 0, 1, request_len, chunk_size_blocks=chunk_size) + add_and_verify_chunked_request(setup, 2, 3, request_len * 2, chunk_size_blocks=chunk_size) + finally: + for worker in setup["ctx_transfer_workers"]: + worker.shutdown() for worker in setup["gen_transfer_workers"]: worker.shutdown() @@ -1496,5 +1959,155 @@ def test_session_has_transferring_tasks_false(): gen_transfer_worker.shutdown() +def add_and_verify_pipelined_request( + setup, + ctx_request_id, + gen_request_id, + request_len, + chunk_size_blocks, + prepopulated_blocks=0, +): + """Pipelined transfer: sender sends chunks incrementally, receiver sends 1. + + ``prepopulated_blocks`` models a context-side prefix-reuse hit. The receiver + still posts one monolithic slice covering the whole prompt (no gen-side reuse), + so the block-data comparison below covers the reused prefix too. + """ + ctx_transfer_workers = setup["ctx_transfer_workers"] + gen_transfer_workers = setup["gen_transfer_workers"] + + ctx_info = _setup_chunked_request(setup, ctx_request_id, gen_request_id, request_len) + ctx_block_ids = ctx_info["ctx_block_ids"] + gen_block_ids = ctx_info["gen_block_ids"] + + sender_sessions = [tw.create_tx_session(ctx_info["ctx_request"]) for tw in ctx_transfer_workers] + for sender_session, block_ids_per_groups in zip(sender_sessions, ctx_block_ids): + _send_prefill_chunks( + block_ids_per_groups, + chunk_size_blocks=chunk_size_blocks, + tokens_per_block=setup["tokens_per_block"], + sender_session=sender_session, + prepopulated_blocks=prepopulated_blocks, + ) + + receiver_sessions = [ + tw.create_rx_session(ctx_info["gen_request"]) for tw in gen_transfer_workers + ] + for recv_session, block_ids_per_groups in zip(receiver_sessions, gen_block_ids): + full_slice = KVSlice( + is_last_slice=True, + block_ids_per_layer_groups=block_ids_per_groups, + ) + recv_session.receive(full_slice) + + for session in sender_sessions: + result = session.wait_complete() + assert result == WaitResult.COMPLETED, f"tx wait_complete returned {result}" + for session in receiver_sessions: + result = session.wait_complete(blocking=True) + assert result == WaitResult.COMPLETED, f"rx wait_complete returned {result}" + + _verify_and_cleanup_chunked(setup, ctx_info, sender_sessions, receiver_sessions) + + +PIPELINED_TEST_CONFIGS = [ + (1, 1, False, 1, 1, False, False, True, "v2_tp1_pp1_pipelined"), + (1, 1, False, 1, 1, False, False, False, "v1_tp1_pp1_pipelined"), +] + + +@pytest.mark.timeout(120) +@pytest.mark.parametrize( + "ctx_tp,ctx_pp,ctx_enable_dp,gen_tp,gen_pp,gen_enable_dp,is_mla,use_v2", + [(c[0], c[1], c[2], c[3], c[4], c[5], c[6], c[7]) for c in PIPELINED_TEST_CONFIGS], + ids=[c[8] for c in PIPELINED_TEST_CONFIGS], +) +def test_transfer_worker_pipelined( + ctx_tp, ctx_pp, ctx_enable_dp, gen_tp, gen_pp, gen_enable_dp, is_mla, use_v2 +): + """Test pipelined transfer: chunks sent incrementally for V1 and V2.""" + if not torch.cuda.is_available(): + pytest.skip("CUDA not available") + tensorrt_llm.logger.set_level("info") + logger.info(f"Test transfer worker {'V2' if use_v2 else 'V1'} with pipelined transfer") + + setup = create_transfer_worker_setup( + ctx_tp=ctx_tp, + ctx_pp=ctx_pp, + ctx_enable_dp=ctx_enable_dp, + gen_tp=gen_tp, + gen_pp=gen_pp, + gen_enable_dp=gen_enable_dp, + is_mla=is_mla, + use_v2=use_v2, + ) + + request_len = setup["request_len"] + tokens_per_block = setup["tokens_per_block"] + total_blocks = (request_len + tokens_per_block - 1) // tokens_per_block + chunk_size = max(1, total_blocks // 2) + + try: + add_and_verify_pipelined_request(setup, 0, 1, request_len, chunk_size_blocks=chunk_size) + add_and_verify_pipelined_request(setup, 2, 3, request_len * 2, chunk_size_blocks=chunk_size) + finally: + for worker in setup["ctx_transfer_workers"]: + worker.shutdown() + for worker in setup["gen_transfer_workers"]: + worker.shutdown() + + +@pytest.mark.timeout(120) +@pytest.mark.parametrize( + "ctx_tp,ctx_pp,ctx_enable_dp,gen_tp,gen_pp,gen_enable_dp,is_mla,use_v2", + [(c[0], c[1], c[2], c[3], c[4], c[5], c[6], c[7]) for c in PIPELINED_TEST_CONFIGS], + ids=[c[8] for c in PIPELINED_TEST_CONFIGS], +) +def test_transfer_worker_pipelined_ctx_prefix_reuse( + ctx_tp, ctx_pp, ctx_enable_dp, gen_tp, gen_pp, gen_enable_dp, is_mla, use_v2 +): + """Pipelined transfer with a ctx-side reuse hit still delivers the reused prefix. + + The scheduler's chunks start after the reused prefix, so without the first-slice + extension in _build_prefill_chunk those blocks are never written and the + generation-side block data diverges. + """ + if not torch.cuda.is_available(): + pytest.skip("CUDA not available") + tensorrt_llm.logger.set_level("info") + + setup = create_transfer_worker_setup( + ctx_tp=ctx_tp, + ctx_pp=ctx_pp, + ctx_enable_dp=ctx_enable_dp, + gen_tp=gen_tp, + gen_pp=gen_pp, + gen_enable_dp=gen_enable_dp, + is_mla=is_mla, + use_v2=use_v2, + ) + + request_len = setup["request_len"] + tokens_per_block = setup["tokens_per_block"] + total_blocks = (request_len + tokens_per_block - 1) // tokens_per_block + chunk_size = max(1, total_blocks // 4) + prepopulated_blocks = max(1, total_blocks // 4) + + try: + add_and_verify_pipelined_request( + setup, + 0, + 1, + request_len, + chunk_size_blocks=chunk_size, + prepopulated_blocks=prepopulated_blocks, + ) + finally: + for worker in setup["ctx_transfer_workers"]: + worker.shutdown() + for worker in setup["gen_transfer_workers"]: + worker.shutdown() + + if __name__ == "__main__": test_transfer_worker_v1(1, 1, False, 1, 1, False, False) diff --git a/tests/unittest/disaggregated/test_transceiver_bounded_polling.py b/tests/unittest/disaggregated/test_transceiver_bounded_polling.py index 76eff6adc439..1bd06d3daa12 100644 --- a/tests/unittest/disaggregated/test_transceiver_bounded_polling.py +++ b/tests/unittest/disaggregated/test_transceiver_bounded_polling.py @@ -127,6 +127,7 @@ def _make_tx_session( session._timeout_s = 0.25 session._need_aux = need_aux session._terminal_status = None + session._exception = None session.receiver_ready = True session.kv_tasks = kv_tasks session.aux_task = aux_task From f6ee37871adf3134e47abf1ee83b84e80a88c6d9 Mon Sep 17 00:00:00 2001 From: Athena Cai Date: Tue, 11 Aug 2026 00:51:11 +0000 Subject: [PATCH 2/3] Address coderabbit comments Signed-off-by: Athena Cai --- .../_torch/disaggregation/base/transfer.py | 7 +- .../_torch/disaggregation/transceiver.py | 17 ++- .../disaggregated/test_chunked_transfer.py | 105 +++++++++++++++++- .../disaggregated/test_kv_transfer.py | 3 + 4 files changed, 124 insertions(+), 8 deletions(-) diff --git a/tensorrt_llm/_torch/disaggregation/base/transfer.py b/tensorrt_llm/_torch/disaggregation/base/transfer.py index cecb3a443086..5103f520a5ab 100644 --- a/tensorrt_llm/_torch/disaggregation/base/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/base/transfer.py @@ -59,7 +59,12 @@ def derive_chunk_block_coords( token_range: Optional[TokenRange], tokens_per_block: int, ) -> tuple[int, int]: - """Derive global chunk block offset and count from a block-aligned token_range.""" + """Derive global chunk block offset and count from a block-aligned token_range. + + Producers of pipelined slices must align every non-final chunk boundary to + ``tokens_per_block``. The final partial block is represented by an aligned + range ending at its enclosing block boundary. + """ if token_range is None: return 0, 0 if tokens_per_block <= 0: diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index 9bae66d90506..0f82ae00b74b 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -602,12 +602,14 @@ def _close_failed_sessions( self, sessions: dict, reqs: dict, failed: list, mark_retired: bool = False ): for rid in failed: - reqs[rid].state = LlmRequestState.DISAGG_TRANS_ERROR - if mark_retired: - reqs[rid].py_kv_send_session_retired = True - sessions[rid].close() - del reqs[rid] - del sessions[rid] + req = reqs.pop(rid, None) + if req is not None: + req.state = LlmRequestState.DISAGG_TRANS_ERROR + if mark_retired: + req.py_kv_send_session_retired = True + session = sessions.pop(rid, None) + if session is not None: + session.close() def _retire_send_session(self, rid: int, req: Optional[LlmRequest] = None) -> None: """Close a send session and bar any later chunk from re-creating it. @@ -732,6 +734,9 @@ def _build_prefill_chunk( chunk_start_block = 0 if is_first_chunk else chunk_start_pos // tpb chunk_end_block = (chunk_end_pos + tpb - 1) // tpb is_last_chunk = req.context_remaining_length == 0 + assert is_last_chunk or chunk_end_pos % tpb == 0, ( + f"non-final prefill chunk end {chunk_end_pos} must be aligned to tokens_per_block={tpb}" + ) prompt_blocks = (req.prompt_len + tpb - 1) // tpb total_blocks = prompt_blocks diff --git a/tests/unittest/disaggregated/test_chunked_transfer.py b/tests/unittest/disaggregated/test_chunked_transfer.py index 43362f75bab0..1b9655549c1b 100644 --- a/tests/unittest/disaggregated/test_chunked_transfer.py +++ b/tests/unittest/disaggregated/test_chunked_transfer.py @@ -660,7 +660,6 @@ def test_pipelined_transfer_disabled_by_default(): transceiver = MagicMock() transceiver._enable_pipelined_transfer = False - transceiver._chunk_size_blocks = 64 result = KvCacheTransceiverV2.pipeline_transfer_enabled.fget(transceiver) assert result is False @@ -886,6 +885,74 @@ def test_pipelined_non_last_chunk_does_not_finalize(): transceiver._finalize_send.assert_not_called() +def test_pipelined_multiple_chunks_use_real_builder_and_tx_session(): + """Drive two chunks through respond_and_send_async and a real TxSession.""" + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 + + rid = 42 + tokens_per_block = 4 + source_block_ids = np.arange(4, dtype=np.int64) + session = TxSession( + request_id=rid, + params=_make_params(rid), + sender=_stub_sender(), + prompt_len=16, + ) + + transceiver = object.__new__(KvCacheTransceiverV2) + transceiver._enable_pipelined_transfer = True + transceiver._send_sessions = {} + transceiver._send_reqs = {} + transceiver._ever_had_send_session = False + transceiver._transfer_worker = SimpleNamespace(create_tx_session=lambda _req: session) + transceiver._reuse_adapter = SimpleNamespace( + tokens_per_block=tokens_per_block, + get_block_ids=lambda _req, _idx, _lg: source_block_ids, + ) + transceiver._page_table = SimpleNamespace( + layer_groups=[SimpleNamespace(sliding_window_size=None)] + ) + transceiver._kv_cache_manager = SimpleNamespace(tokens_per_block=tokens_per_block) + transceiver._dp_rank = 0 + transceiver._context_info_endpoint = "ctx" + + request = SimpleNamespace( + py_disaggregated_params=_make_params(rid), + request_id=rid, + py_request_id=rid, + prompt_len=16, + py_beam_width=1, + py_kv_send_session_retired=False, + prepopulated_prompt_len=0, + is_generation_only_request=lambda: False, + set_kv_cache_transfer_start=lambda _ts: None, + state=LlmRequestState.CONTEXT_INIT, + ) + + request.py_last_context_chunk = (0, 8) + request.context_remaining_length = 8 + transceiver.respond_and_send_async(request) + + request.py_last_context_chunk = (8, 16) + request.context_remaining_length = 0 + transceiver.respond_and_send_async(request) + + assert [task._slice.token_range for task in session.kv_tasks] == [ + TokenRange(start=0, end=8), + TokenRange(start=8, end=16), + ] + assert [task._slice.block_ids_per_layer_groups[0].tolist() for task in session.kv_tasks] == [ + [0, 1], + [2, 3], + ] + assert [task._slice.is_last_slice for task in session.kv_tasks] == [False, True] + assert transceiver._send_sessions == {rid: session} + assert transceiver._send_reqs == {rid: request} + assert request.state == LlmRequestState.DISAGG_CONTEXT_TRANS_IN_PROGRESS + + session.close() + + # --------------------------------------------------------------------------- # Retired send sessions # --------------------------------------------------------------------------- @@ -906,6 +973,23 @@ def _make_retirable_request(retired: bool, rid: int = 42): ) +def test_close_failed_send_session_without_request(): + """A failed session must be retired even before its request is registered.""" + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 + + session = MagicMock() + sessions = {42: session} + reqs = {} + + KvCacheTransceiverV2._close_failed_sessions( + MagicMock(), sessions, reqs, failed=[42], mark_retired=True + ) + + session.close.assert_called_once_with() + assert sessions == {} + assert reqs == {} + + def test_get_or_create_send_session_refuses_retired_request(): """Closing a send session drops the peer registration, so a new one is inert. @@ -1026,6 +1110,25 @@ def _build_prefill_chunk_for( return KvCacheTransceiverV2._build_prefill_chunk(transceiver, req) +def test_build_prefill_chunk_rejects_unaligned_non_final_end(): + """Reject a partial destination block that the next chunk would overlap.""" + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 + + transceiver = MagicMock() + transceiver._kv_cache_manager.tokens_per_block = _REUSE_TPB + transceiver._send_reqs = {} + + req = MagicMock() + req.py_disaggregated_params = DisaggregatedParams(disagg_request_id=42) + req.py_beam_width = 1 + req.prepopulated_prompt_len = 0 + req.py_last_context_chunk = (0, 6) + req.context_remaining_length = 10 + + with pytest.raises(AssertionError, match="non-final prefill chunk end 6 must be aligned"): + KvCacheTransceiverV2._build_prefill_chunk(transceiver, req) + + def test_first_chunk_covers_ctx_prefix_reuse(): """The reused prefix is resident but no chunk spans it, so slice 0 extends to block 0.""" kv_slice = _build_prefill_chunk_for( diff --git a/tests/unittest/disaggregated/test_kv_transfer.py b/tests/unittest/disaggregated/test_kv_transfer.py index 6039fbf1b8da..30d4396184b0 100644 --- a/tests/unittest/disaggregated/test_kv_transfer.py +++ b/tests/unittest/disaggregated/test_kv_transfer.py @@ -1997,6 +1997,7 @@ def add_and_verify_pipelined_request( full_slice = KVSlice( is_last_slice=True, block_ids_per_layer_groups=block_ids_per_groups, + token_range=TokenRange(start=0, end=request_len), ) recv_session.receive(full_slice) @@ -2013,6 +2014,8 @@ def add_and_verify_pipelined_request( PIPELINED_TEST_CONFIGS = [ (1, 1, False, 1, 1, False, False, True, "v2_tp1_pp1_pipelined"), (1, 1, False, 1, 1, False, False, False, "v1_tp1_pp1_pipelined"), + (2, 1, False, 2, 1, False, False, True, "v2_tp2_pp1_pipelined"), + (2, 1, False, 2, 1, False, False, False, "v1_tp2_pp1_pipelined"), ] From 456afe2b0fc2e3e514a381e60b2e865e44814b60 Mon Sep 17 00:00:00 2001 From: Athena Cai Date: Tue, 11 Aug 2026 22:21:23 +0000 Subject: [PATCH 3/3] Add bounded tail block-ID retrieval for chunked KV transfer _build_prefill_chunk now asks each layer group only for the block range the current chunk needs instead of rebuilding a whole-request slice and projecting it. A windowed group's range starts at the request's final-window stale boundary: the sender worker floors src_start there without dropping the blocks under it, so a run reaching further back would be paired with the window tail's destination blocks. A chunk entirely below the window skips the cache query altogether. V1 gains a bounded C++/nanobind query and keeps the block-ID to pool-slot translation get_block_ids does. V2 cuts the range out of the aggregated page list, since neither of its backends has a bounded query. Signed-off-by: Athena Cai --- .../batch_manager/kvCacheManager.h | 9 + .../batch_manager/kvCacheManager.cpp | 29 +++ .../nanobind/batch_manager/kvCacheManager.cpp | 5 + .../batch_manager/kvCacheManagerTest.cpp | 60 +++++ .../disaggregation/resource/cache_reuse.py | 75 ++++++ .../_torch/disaggregation/transceiver.py | 118 ++++++---- .../_torch/pyexecutor/resource_manager.py | 57 ++++- .../disaggregated/test_cache_reuse_adapter.py | 138 +++++++++++ .../disaggregated/test_chunked_transfer.py | 22 +- .../disaggregated/test_kv_transfer.py | 218 ++++++++++++++++-- 10 files changed, 654 insertions(+), 77 deletions(-) diff --git a/cpp/include/tensorrt_llm/batch_manager/kvCacheManager.h b/cpp/include/tensorrt_llm/batch_manager/kvCacheManager.h index 4ce5fa94b131..7e3ed20ea1af 100644 --- a/cpp/include/tensorrt_llm/batch_manager/kvCacheManager.h +++ b/cpp/include/tensorrt_llm/batch_manager/kvCacheManager.h @@ -2170,6 +2170,15 @@ class BaseKVCacheManager std::vector const& requestIds, SizeType32 windowSize) const = 0; + //! \brief Get resident block ids in the absolute request block range [blockBegin, blockEnd), per beam. + //! \details This non-virtual convenience wrapper copies only the requested range out of the sequence's persistent + //! block table, dropping front-detached SWA blocks (which keep their recycled ids in the raw table). The result is + //! always the *contiguous* run ending at blockEnd, i.e. ordinals [blockEnd - result.size(), blockEnd), so a caller + //! can recover each id's block ordinal from blockEnd and the result size alone. Requesting a blockEnd past the + //! sequence's allocated blocks would break that guarantee and is rejected. + [[nodiscard]] std::vector> getCacheBlockIdsRange( + LlmRequest::RequestIdType requestId, SizeType32 windowSize, SizeType32 blockBegin, SizeType32 blockEnd) const; + /// @brief Get the last block id (beam 0) for a given sequence and window size [[nodiscard]] virtual std::optional getLastBlockId(LlmRequest::RequestIdType requestId) const = 0; diff --git a/cpp/tensorrt_llm/batch_manager/kvCacheManager.cpp b/cpp/tensorrt_llm/batch_manager/kvCacheManager.cpp index 18a8d83e828c..2b6c016040b1 100644 --- a/cpp/tensorrt_llm/batch_manager/kvCacheManager.cpp +++ b/cpp/tensorrt_llm/batch_manager/kvCacheManager.cpp @@ -38,6 +38,7 @@ #include "tensorrt_llm/runtime/worldConfig.h" #include +#include #include #include #include @@ -4629,6 +4630,34 @@ std::vector> const& KVCacheManager::getCacheBlockIds( return getSequence(requestId).getCacheBlockIds(windowSize); } +std::vector> BaseKVCacheManager::getCacheBlockIdsRange( + LlmRequest::RequestIdType requestId, SizeType32 windowSize, SizeType32 blockBegin, SizeType32 blockEnd) const +{ + TLLM_CHECK_WITH_INFO(blockBegin >= 0, "blockBegin must be non-negative"); + TLLM_CHECK_WITH_INFO(blockEnd >= 0, "blockEnd must be non-negative"); + TLLM_CHECK_WITH_INFO(blockBegin <= blockEnd, "blockBegin must not exceed blockEnd"); + // Read the block table and the eviction count off the same sequence, so a manager that overrides one but not the + // other cannot hand back a block table and a front-eviction count that disagree. + auto const& sequence = getSequence(requestId); + auto const& blockIdsPerBeam = sequence.getCacheBlockIds(windowSize); + auto const firstResidentBlock = static_cast(sequence.getNumFrontBlocksRemoved(windowSize)); + auto const end = static_cast(blockEnd); + std::vector> result; + result.reserve(blockIdsPerBeam.size()); + for (auto const& blockIds : blockIdsPerBeam) + { + TLLM_CHECK_WITH_INFO(end <= blockIds.size(), + "blockEnd=%d exceeds the %zu blocks allocated for request %lu at windowSize=%d; the result would not end " + "at blockEnd, and callers recover block ordinals from its size", + blockEnd, blockIds.size(), static_cast(requestId), windowSize); + auto const begin = std::min(end, std::max(firstResidentBlock, static_cast(blockBegin))); + auto const beginOffset = static_cast(begin); + auto const endOffset = static_cast(end); + result.emplace_back(blockIds.begin() + beginOffset, blockIds.begin() + endOffset); + } + return result; +} + std::vector KVCacheManager::commitAndGetBlockHashesForRequest( LlmRequest const& llmRequest, SizeType32 windowSize) { diff --git a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManager.cpp b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManager.cpp index 559934e1ef72..7ce416157cf6 100644 --- a/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManager.cpp +++ b/cpp/tensorrt_llm/nanobind/batch_manager/kvCacheManager.cpp @@ -649,6 +649,11 @@ void tb::kv_cache_manager::KVCacheManagerBindings::initBindings(nb::module_& m) nb::arg("request_id"), nb::arg("window_size"), nb::call_guard()) .def("get_batch_cache_block_ids", &BaseKVCacheManager::getBatchCacheBlockIds, nb::call_guard()) + // Deliberately keeps the GIL, unlike its neighbours: it reaches the virtual getSequence, whose trampoline can + // call back into a Python subclass. The copy it makes is bounded by the requested range, so there is little + // to gain from releasing. + .def("get_cache_block_ids_range", &BaseKVCacheManager::getCacheBlockIdsRange, nb::arg("request_id"), + nb::arg("window_size"), nb::arg("block_begin"), nb::arg("block_end")) .def("flush_iteration_events", &BaseKVCacheManager::flushIterationEvents, nb::call_guard()) .def("sync_transfer_manager_with_buffer_manager", &BaseKVCacheManager::syncTransferManagerWithBufferManager, diff --git a/cpp/tests/unit_tests/batch_manager/kvCacheManagerTest.cpp b/cpp/tests/unit_tests/batch_manager/kvCacheManagerTest.cpp index 764d3236c609..88e045f82c9a 100644 --- a/cpp/tests/unit_tests/batch_manager/kvCacheManagerTest.cpp +++ b/cpp/tests/unit_tests/batch_manager/kvCacheManagerTest.cpp @@ -9703,6 +9703,66 @@ TEST_F(KVCacheManagerTest, VSWABlockStoredDuringGeneration) EXPECT_EQ(blockManager.getNumFreeBlocks(), blocksInPrimaryPool); } +// getCacheBlockIdsRange must return the contiguous run of resident blocks ending at blockEnd. +// detachFrontBlock leaves the recycled physical id of an out-of-window block in the raw block +// table, so a range starting before the eviction count must skip it: the pipelined KV transceiver +// would otherwise read KV out of a block that already belongs to another request. +TEST_F(KVCacheManagerTest, VSWAGetCacheBlockIdsRangeExcludesDetachedFrontBlocks) +{ + auto constexpr blocksInPrimaryPool = 10; + auto const stream = std::make_shared(); + tr::SamplingConfig const samplingConfig{kVSWA_BEAM_WIDTH}; + auto kvCacheManager = makeVSWAManager(blocksInPrimaryPool, /*enableBlockReuse=*/true, stream); + + // Seq 0: 11 input tokens covering B0=[1000..1003], B1=[1004..1007], B2=[1008..1010] (partial). + auto inputTokens0 = std::make_shared(11); + std::iota(inputTokens0->begin(), inputTokens0->end(), kVSWA_FIRST_TOKEN); + auto llmRequest0 + = std::make_shared(0, kVSWA_MAX_NEW_TOKENS, inputTokens0, samplingConfig, kVSWA_IS_STREAMING); + addSequenceForTest(*kvCacheManager, 0, 11, kVSWA_BEAM_WIDTH, llmRequest0); + tensorrt_llm::testing::KvCacheManagerTestUtil::simulatePrefillCompletion(*llmRequest0); + kvCacheManager->storeContextBlocks(*llmRequest0); + + auto const& rawBlockIds = kvCacheManager->getCacheBlockIds(0, kVSWA_ATTENTION_WINDOW).at(kVSWA_BEAM_IDX); + ASSERT_EQ(rawBlockIds.size(), 3U); + + // Before any eviction the range is exactly what was asked for. + EXPECT_EQ(kvCacheManager->getSequence(0).getNumFrontBlocksRemoved(kVSWA_ATTENTION_WINDOW), 0); + EXPECT_THAT(kvCacheManager->getCacheBlockIdsRange(0, kVSWA_ATTENTION_WINDOW, 0, 3).at(kVSWA_BEAM_IDX), + ::testing::ElementsAre(rawBlockIds.at(0), rawBlockIds.at(1), rawBlockIds.at(2))); + EXPECT_THAT(kvCacheManager->getCacheBlockIdsRange(0, kVSWA_ATTENTION_WINDOW, 1, 3).at(kVSWA_BEAM_IDX), + ::testing::ElementsAre(rawBlockIds.at(1), rawBlockIds.at(2))); + // One entry per beam, even when the range is empty. + EXPECT_EQ(kvCacheManager->getCacheBlockIdsRange(0, kVSWA_ATTENTION_WINDOW, 2, 2).size(), + static_cast(kVSWA_BEAM_WIDTH)); + EXPECT_THAT(kvCacheManager->getCacheBlockIdsRange(0, kVSWA_ATTENTION_WINDOW, 2, 2).at(kVSWA_BEAM_IDX), + ::testing::IsEmpty()); + + // Generation step: numTokens becomes 12; adjustBlocksIfNeeded detaches B0 (12-0*4 >= 8+4). + llmRequest0->addNewToken(kVSWA_FIRST_TOKEN + 11, kVSWA_BEAM_IDX); + kvCacheManager->addToken(0); + ASSERT_EQ(kvCacheManager->getSequence(0).getNumFrontBlocksRemoved(kVSWA_ATTENTION_WINDOW), 1); + // B0's recycled id is still in the raw table, so the range query is the only thing standing + // between the transceiver and another request's data. + ASSERT_EQ(kvCacheManager->getCacheBlockIds(0, kVSWA_ATTENTION_WINDOW).at(kVSWA_BEAM_IDX).size(), 3U); + + EXPECT_THAT(kvCacheManager->getCacheBlockIdsRange(0, kVSWA_ATTENTION_WINDOW, 0, 1).at(kVSWA_BEAM_IDX), + ::testing::IsEmpty()); + EXPECT_THAT(kvCacheManager->getCacheBlockIdsRange(0, kVSWA_ATTENTION_WINDOW, 0, 3).at(kVSWA_BEAM_IDX), + ::testing::ElementsAre(rawBlockIds.at(1), rawBlockIds.at(2))); + EXPECT_THAT(kvCacheManager->getCacheBlockIdsRange(0, kVSWA_ATTENTION_WINDOW, 2, 3).at(kVSWA_BEAM_IDX), + ::testing::ElementsAre(rawBlockIds.at(2))); + + // Bad bounds are programming errors, and so is reading past the allocated blocks: the result + // would no longer end at blockEnd, which is how callers recover each id's block ordinal. + EXPECT_THROW((void) kvCacheManager->getCacheBlockIdsRange(0, kVSWA_ATTENTION_WINDOW, -1, 2), std::runtime_error); + EXPECT_THROW((void) kvCacheManager->getCacheBlockIdsRange(0, kVSWA_ATTENTION_WINDOW, 0, -1), std::runtime_error); + EXPECT_THROW((void) kvCacheManager->getCacheBlockIdsRange(0, kVSWA_ATTENTION_WINDOW, 2, 1), std::runtime_error); + EXPECT_THROW((void) kvCacheManager->getCacheBlockIdsRange(0, kVSWA_ATTENTION_WINDOW, 0, 4), std::runtime_error); + + EXPECT_NO_THROW(static_cast(kvCacheManager->removeSequence(0, llmRequest0))); +} + // Verify that when an OOW block is stolen by another sequence, storeBlocks does // not restore that missing anchor under the original sequence's key or corrupt // the acquiring sequence's trie, and all blocks are properly released. diff --git a/tensorrt_llm/_torch/disaggregation/resource/cache_reuse.py b/tensorrt_llm/_torch/disaggregation/resource/cache_reuse.py index 142161093f12..1dae7d7a9cfa 100644 --- a/tensorrt_llm/_torch/disaggregation/resource/cache_reuse.py +++ b/tensorrt_llm/_torch/disaggregation/resource/cache_reuse.py @@ -22,6 +22,7 @@ from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2 from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm.runtime.kv_cache_manager_v2 import BAD_PAGE_INDEX from .page import AttentionLayerGroup from .utils import get_global_layer_ids @@ -74,6 +75,35 @@ def get_block_ids( must translate before returning. """ + @abstractmethod + def get_block_ids_range( + self, + req: LlmRequest, + group_idx: int, + lg: AttentionLayerGroup, + block_begin: int, + block_end: int, + ) -> np.ndarray: + """Resident block IDs for absolute request ordinals ``[block_begin, block_end)``. + + The result is the *contiguous* run of resident blocks ending at + ``block_end``, i.e. ordinals ``[block_end - len(result), block_end)``. + It can be shorter than the requested range when leading blocks have been + evicted, as under sliding-window attention; blocks at or before any gap + are dropped rather than compacted. + + Callers must not reorder or reinterpret the result positionally beyond + that rule: the receiver reconstructs each layer group's starting token + from ``block_end`` and the result length, so a result that is not a + contiguous run ending at ``block_end`` would be written to the wrong + offsets. + + Empty ranges are valid. Negative or reversed bounds raise + ``ValueError``. A ``block_end`` past the request's allocated blocks also + raises, though the exception type is backend-specific (``ValueError`` + from V2, ``RuntimeError`` out of the C++ manager for V1). + """ + @abstractmethod def commit_blocks_for_reuse(self, req: LlmRequest) -> None: """Commit KV blocks to radix tree for future prefix reuse. @@ -123,6 +153,23 @@ def get_block_ids(self, req, group_idx, lg): # noqa: ARG002 ) return np.asarray(pool_indices, dtype=np.int64) + def get_block_ids_range(self, req, group_idx, lg, block_begin, block_end): # noqa: ARG002 + window_size = lg.sliding_window_size + # V1 layer groups always carry the manager's window key; see get_block_ids. + assert window_size is not None + raw_ids = self._mgr.get_cache_indices_range( + req.py_request_id, + block_begin=block_begin, + block_end=block_end, + window_size=window_size, + ) + if not raw_ids: + return np.array([], dtype=np.int64) + # Same block_id -> primary-pool slot translation get_block_ids does; the + # two diverge once host offload is enabled. + pool_indices = self._mgr.get_memory_pool_block_indices(raw_ids, window_size=window_size) + return np.asarray(pool_indices, dtype=np.int64) + def commit_blocks_for_reuse(self, req: LlmRequest) -> None: if not self.enable_block_reuse: return @@ -165,6 +212,34 @@ def get_block_ids(self, req, group_idx, lg): # noqa: ARG002 dtype=np.int64, ) + def get_block_ids_range(self, req, group_idx, lg, block_begin, block_end): # noqa: ARG002 + if block_begin < 0 or block_end < 0: + raise ValueError("block range bounds must be non-negative") + if block_begin > block_end: + raise ValueError("block_begin must not exceed block_end") + # Neither V2 backend has a bounded query, so read the whole aggregated + # list -- keeping the placeholders, which is what makes an entry's index + # its block ordinal -- and cut the range out of it here. + all_block_ids = np.fromiter( + self._mgr.kv_cache_map[req.py_request_id].get_aggregated_page_indices( + group_idx, valid_only=False + ), + dtype=np.int64, + ) + if block_end > all_block_ids.size: + raise ValueError( + f"block_end={block_end} exceeds the {all_block_ids.size} allocated blocks; " + "the result would not end at block_end, and callers recover block " + "ordinals from its length" + ) + block_ids = all_block_ids[block_begin:block_end] + # Drop everything up to the last gap rather than compacting around it: + # a life cycle that keeps sink tokens resident has a sink prefix plus a + # window suffix, and returning both would misreport where the suffix + # starts. + gaps = np.flatnonzero(block_ids == BAD_PAGE_INDEX) + return block_ids[gaps[-1] + 1 :] if gaps.size else block_ids + def commit_blocks_for_reuse(self, req: LlmRequest) -> None: self._mgr.try_commit_blocks(req) diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index 0f82ae00b74b..aa6640fdc329 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -32,7 +32,6 @@ TxSessionBase, WaitResult, get_unique_rid, - project_blocks_to_global_chunk, ) from tensorrt_llm._torch.disaggregation.native.bounce import ( config_from_size as bounce_config_from_size, @@ -60,6 +59,9 @@ from tensorrt_llm.llmapi.llm_args import CacheTransceiverConfig from tensorrt_llm.mapping import Mapping +_EMPTY_BLOCK_IDS = np.array([], dtype=np.int64) +_EMPTY_BLOCK_IDS.flags.writeable = False + def _find_consensus_request_ids(request_ids_all_ranks, sync_size): frequency_map = defaultdict(int) @@ -252,30 +254,20 @@ def __enter__(self): def __exit__(self, _exc_type, _exc_val, _exc_tb): self.shutdown() - def _create_kv_slice( - self, - req: LlmRequest, - resident_block_end: Optional[int] = None, - ) -> KVSlice: - """Create a KV slice from the source's currently resident blocks. + def _create_kv_slice(self, req: LlmRequest) -> KVSlice: + """Create a KV slice covering the whole prompt. + + Pipelined prefill builds its per-chunk slices in ``_build_prefill_chunk`` + instead, which bounds each layer group's block range directly. Args: req: Request whose KV blocks are being described. - resident_block_end: Exclusive logical block boundary to include. - Pipelined prefill passes the current chunk end to exclude - full-prompt blocks that V1 reserved but has not computed yet. - ``None`` includes the complete prompt. """ adapter = self._reuse_adapter tpb = adapter.tokens_per_block assert self._page_table is not None layer_groups = self._page_table.layer_groups - prompt_blocks = (req.prompt_len + tpb - 1) // tpb - resident_blocks = ( - prompt_blocks - if resident_block_end is None - else min(max(0, resident_block_end), prompt_blocks) - ) + resident_blocks = (req.prompt_len + tpb - 1) // tpb is_gen_only = req.is_generation_only_request() cached_per_lg = ( @@ -331,22 +323,22 @@ def _create_kv_slice( groups.append(block_ids) - mamba_state_index = None - if isinstance(self._kv_cache_manager, MambaHybridCacheManagerV2): - if self._kv_cache_manager.local_num_mamba_layers > 0: - mamba_state_index = self._kv_cache_manager._request_id_to_state_index[ - req.py_request_id - ] - elif isinstance(self._kv_cache_manager, MambaHybridCacheManager): - mamba_state_index = self._kv_cache_manager.mamba_cache_index[req.py_request_id] - return KVSlice( is_last_slice=True, block_ids_per_layer_groups=groups, - mamba_state_index=mamba_state_index, + mamba_state_index=self._get_mamba_state_index(req), token_range=token_range, ) + def _get_mamba_state_index(self, req: LlmRequest) -> Optional[int]: + if isinstance(self._kv_cache_manager, MambaHybridCacheManagerV2): + if self._kv_cache_manager.local_num_mamba_layers > 0: + return self._kv_cache_manager._request_id_to_state_index[req.py_request_id] + return None + if isinstance(self._kv_cache_manager, MambaHybridCacheManager): + return self._kv_cache_manager.mamba_cache_index[req.py_request_id] + return None + def _slice_num_bytes(self, slice: KVSlice) -> int: """Local-rank KV bytes covered by a slice (sum of num_valid_blocks * pool.slot_bytes), enough to populate kv_cache_size and unblock the perf-metric timestamps that gate on it.""" @@ -738,26 +730,68 @@ def _build_prefill_chunk( f"non-final prefill chunk end {chunk_end_pos} must be aligned to tokens_per_block={tpb}" ) + # Keep the full prompt span for destination projection while requesting + # only the current absolute block range from each source layer group. prompt_blocks = (req.prompt_len + tpb - 1) // tpb total_blocks = prompt_blocks chunk_start = min(chunk_start_block, total_blocks) chunk_end = min(chunk_end_block, total_blocks) chunk_block_count = max(0, chunk_end - chunk_start) - # V1 reserves the full prompt up front, while V2 grows its source list - # incrementally. Normalize both to the current chunk boundary before - # projecting so full-prompt SWA pages are never treated as current pages. - base_slice = self._create_kv_slice(req, resident_block_end=chunk_end) - all_block_ids = base_slice.block_ids_per_layer_groups - chunk_block_ids = [ - project_blocks_to_global_chunk( - block_ids, - chunk_block_offset=chunk_start, - chunk_block_count=chunk_block_count, - resident_block_end=chunk_end, - ) - for block_ids in all_block_ids - ] + assert self._page_table is not None + layer_groups = self._page_table.layer_groups + if chunk_block_count == 0: + chunk_block_ids = [_EMPTY_BLOCK_IDS] * len(layer_groups) + else: + chunk_block_ids = [] + for group_idx, layer_group in enumerate(layer_groups): + block_ids = _EMPTY_BLOCK_IDS + if not isinstance(layer_group, MambaLayerGroup): + range_begin = chunk_start + if layer_group.sliding_window_size is not None: + # The stale boundary is a property of the whole request, + # so it uses prompt_len rather than this chunk's end. A + # windowed group's run must not reach below it: + # _build_kv_write_meta raises src_start to + # stale_end * tpb without dropping the blocks under it, + # which would pair this run's head with the window + # tail's destination blocks. A chunk entirely below the + # window contributes nothing, so the cache is not + # queried at all — the common case for a short window + # and a long prompt. + stale_end = max( + 0, (req.prompt_len + 1 - layer_group.sliding_window_size) // tpb + ) + range_begin = max(range_begin, stale_end) + range_block_count = chunk_end - range_begin + if range_block_count > 0: + block_ids = self._reuse_adapter.get_block_ids_range( + req, + group_idx, + layer_group, + block_begin=range_begin, + block_end=chunk_end, + ) + # The receiver derives each group's starting token from + # chunk_end minus the number of blocks it received, so a + # group may only ever drop *leading* blocks. Anything + # else writes KV to the wrong offsets, silently. + if block_ids.size > range_block_count: + raise ValueError( + f"layer group {group_idx} returned {block_ids.size} blocks for " + f"range [{range_begin}, {chunk_end}), which spans only " + f"{range_block_count}" + ) + if ( + layer_group.sliding_window_size is None + and block_ids.size != chunk_block_count + ): + raise ValueError( + f"chunk [{chunk_start}, {chunk_end}) of layer group {group_idx} " + f"is not fully resident: got {block_ids.size} of " + f"{chunk_block_count} blocks" + ) + chunk_block_ids.append(block_ids) chunk_token_range = None if chunk_block_count > 0: chunk_token_range = TokenRange( @@ -768,7 +802,7 @@ def _build_prefill_chunk( return KVSlice( is_last_slice=is_last_chunk, block_ids_per_layer_groups=chunk_block_ids, - mamba_state_index=base_slice.mamba_state_index, + mamba_state_index=self._get_mamba_state_index(req), token_range=chunk_token_range, total_blocks=total_blocks, ) diff --git a/tensorrt_llm/_torch/pyexecutor/resource_manager.py b/tensorrt_llm/_torch/pyexecutor/resource_manager.py index 153fe93e82ae..c11775fab67c 100644 --- a/tensorrt_llm/_torch/pyexecutor/resource_manager.py +++ b/tensorrt_llm/_torch/pyexecutor/resource_manager.py @@ -599,6 +599,13 @@ def append_to_kv_heads_per_layer(num_kv_heads_per_layer: List[int], f"Adjusted attention window size to {self.max_seq_len} in blocks_per_window" ) + # Cache the layer-to-window mapping now that window clamping is done: + # block-ID queries run on the executor's per-iteration path and must not + # rebuild it. Note this is derived from max_attention_window_vec, which + # the non-SELF branch above leaves at its pre-rewrite value. + self._window_size_by_layer_offset = self._get_layer_offset_to_window_size( + ) + # Use the provided execution stream for proper synchronization with KVCacheTransferManager. # The execution stream is the stream where model forward kernels run, and KVCacheTransferManager # needs to synchronize with it for onboard/offload operations. @@ -1443,6 +1450,15 @@ def get_memory_pool_block_indices(self, block_ids: List[int], *, return self.impl.get_memory_pool_block_indices(list(block_ids), window_size) + def _resolve_cache_window_size(self, layer_idx: Optional[int], + window_size: Optional[int]) -> int: + if window_size is not None or layer_idx is None: + return self._resolve_window_size( + window_size, + "layer_idx or window_size must be provided for VSWA") + layer_offset = self.layer_offsets[layer_idx] + return self._window_size_by_layer_offset[layer_offset] + def get_batch_cache_indices( self, request_ids: List[int], @@ -1452,17 +1468,7 @@ def get_batch_cache_indices( num_blocks_per_seq: Optional[Sequence[int]] = None, ) -> List[List[int]]: beam_width = beam_width or 1 - if window_size is None: - if layer_idx is None: - window_size = self._resolve_window_size( - window_size, - "layer_idx or window_size must be provided for VSWA") - else: - layer_offset = self.layer_offsets[layer_idx] - # Explicit layer_offset -> window_size mapping (no modulo - # masking length mismatches between pattern and num_local_layers). - window_size = self._get_layer_offset_to_window_size( - )[layer_offset] + window_size = self._resolve_cache_window_size(layer_idx, window_size) result = self.impl.get_batch_cache_block_ids(request_ids, window_size) for i in range(len(result)): @@ -1476,6 +1482,35 @@ def get_batch_cache_indices( result[i] = result[i][:num_blocks_per_seq[i]] return result + def get_cache_indices_range( + self, + request_id: int, + block_begin: int, + block_end: int, + layer_idx: Optional[int] = None, + window_size: Optional[int] = None, + ) -> List[int]: + """Return resident cache indices in absolute range ``[block_begin, block_end)``. + + The result is the contiguous run of resident blocks ending at + ``block_end``; front-detached SWA blocks are excluded. See + ``BaseKVCacheManager::getCacheBlockIdsRange`` for the full contract. + """ + # Mirror KVCacheManagerV2, which raises ValueError for these, so both + # backends of CacheReuseAdapter.get_block_ids_range agree. + if block_begin < 0 or block_end < 0: + raise ValueError("block range bounds must be non-negative") + if block_begin > block_end: + raise ValueError("block_begin must not exceed block_end") + window_size = self._resolve_cache_window_size(layer_idx, window_size) + per_beam = self.impl.get_cache_block_ids_range(request_id, window_size, + block_begin, block_end) + if len(per_beam) != 1: + raise ValueError( + f"Chunked/pipelined KV transfer requires beam_width=1, got {len(per_beam)} beams" + ) + return list(per_beam[0]) + def get_batch_cache_indices_flat( self, request_ids: List[int], diff --git a/tests/unittest/disaggregated/test_cache_reuse_adapter.py b/tests/unittest/disaggregated/test_cache_reuse_adapter.py index d9f8e7210f8f..9397f550fe1e 100644 --- a/tests/unittest/disaggregated/test_cache_reuse_adapter.py +++ b/tests/unittest/disaggregated/test_cache_reuse_adapter.py @@ -25,10 +25,12 @@ from tensorrt_llm._torch.disaggregation.resource.cache_reuse import ( CacheReuseAdapter, _CacheReuseAdapterV1, + _CacheReuseAdapterV2, ) from tensorrt_llm._torch.disaggregation.resource.page import AttentionLayerGroup, LocalLayer from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm.runtime.kv_cache_manager_v2 import BAD_PAGE_INDEX pytestmark = pytest.mark.cpu_only @@ -110,6 +112,139 @@ def test_dst_extra_draft_block(self): # --------------------------------------------------------------------------- +class TestBlockRangeAdapters: + def test_v1_adapter_requests_only_block_range(self): + """Only the requested range is read, and its ids reach the pool translation.""" + mgr = MagicMock() + mgr.get_cache_indices_range.return_value = [13, 14, 15] + mgr.get_memory_pool_block_indices.return_value = [113, 114, 115] + req = SimpleNamespace(py_request_id=1) + + block_ids = _CacheReuseAdapterV1(mgr).get_block_ids_range( + req, 0, _lg(window=128), block_begin=13, block_end=16 + ) + + mgr.get_cache_indices_range.assert_called_once_with( + 1, block_begin=13, block_end=16, window_size=128 + ) + mgr.get_memory_pool_block_indices.assert_called_once_with([13, 14, 15], window_size=128) + np.testing.assert_array_equal(block_ids, [113, 114, 115]) + + def test_v1_adapter_skips_translation_for_empty_range(self): + mgr = MagicMock() + mgr.get_cache_indices_range.return_value = [] + req = SimpleNamespace(py_request_id=1) + + block_ids = _CacheReuseAdapterV1(mgr).get_block_ids_range( + req, 0, _lg(window=128), block_begin=16, block_end=16 + ) + + assert block_ids.size == 0 + mgr.get_memory_pool_block_indices.assert_not_called() + + @staticmethod + def _v2_adapter(*slot_ids): + """Adapter over a cache whose aggregated list is ``slot_ids``. + + ``BAD_PAGE_INDEX`` marks a block that is not resident, which is what + ``get_aggregated_page_indices`` yields with ``valid_only=False``. + """ + kv_cache = MagicMock() + kv_cache.get_aggregated_page_indices.return_value = list(slot_ids) + mgr = MagicMock() + mgr.kv_cache_map = {1: kv_cache} + return _CacheReuseAdapterV2(mgr), kv_cache + + def _v2_range(self, *slot_ids, block_begin, block_end): + adapter, kv_cache = self._v2_adapter(*slot_ids) + block_ids = adapter.get_block_ids_range( + SimpleNamespace(py_request_id=1), + 2, + _lg(), + block_begin=block_begin, + block_end=block_end, + ) + kv_cache.get_aggregated_page_indices.assert_called_once_with(2, valid_only=False) + return block_ids.tolist() + + def test_v2_range_honors_block_begin(self): + assert self._v2_range(10, 11, 12, 13, block_begin=2, block_end=4) == [12, 13] + + def test_v2_range_drops_leading_gap(self): + """A front-evicted prefix shortens the run without shifting the rest.""" + BAD = BAD_PAGE_INDEX + assert self._v2_range(BAD, BAD, 12, 13, block_begin=0, block_end=4) == [12, 13] + + def test_v2_range_drops_everything_before_an_interior_gap(self): + """Sink blocks before a stale hole must not be compacted onto the window suffix. + + Returning [10, 12, 13] here would make the receiver read the run as + ordinals 1..3 and write block 10's KV over ordinal 1. + """ + assert self._v2_range(10, BAD_PAGE_INDEX, 12, 13, block_begin=0, block_end=4) == [12, 13] + + def test_v2_range_is_empty_when_the_last_block_is_absent(self): + assert self._v2_range(10, 11, BAD_PAGE_INDEX, block_begin=0, block_end=3) == [] + + def test_v2_range_ignores_gaps_outside_the_range(self): + """A block past block_end cannot shorten the run.""" + assert self._v2_range(10, 11, BAD_PAGE_INDEX, block_begin=0, block_end=2) == [10, 11] + + @pytest.mark.parametrize("block_begin,block_end", [(-1, 1), (0, -1), (2, 1)]) + def test_v2_range_rejects_invalid_bounds(self, block_begin, block_end): + adapter, kv_cache = self._v2_adapter() + + with pytest.raises(ValueError): + adapter.get_block_ids_range( + SimpleNamespace(py_request_id=1), + 0, + _lg(), + block_begin=block_begin, + block_end=block_end, + ) + + kv_cache.get_aggregated_page_indices.assert_not_called() + + def test_v2_range_rejects_block_end_past_allocation(self): + with pytest.raises(ValueError, match="exceeds the 2 allocated blocks"): + self._v2_range(10, 11, block_begin=0, block_end=3) + + @pytest.mark.parametrize("block_begin,block_end", [(-1, 1), (0, -1), (2, 1)]) + def test_v1_manager_rejects_invalid_bounds(self, block_begin, block_end): + """V1 must reject the same bounds as V2, with the same exception type.""" + manager = object.__new__(KVCacheManager) + manager.layer_offsets = [0] + manager._window_size_by_layer_offset = {0: 128} + manager.impl = MagicMock() + + with pytest.raises(ValueError): + manager.get_cache_indices_range(1, block_begin, block_end, layer_idx=0) + + manager.impl.get_cache_block_ids_range.assert_not_called() + + def test_v1_manager_resolves_layer_window_once_and_forwards_range(self): + manager = object.__new__(KVCacheManager) + manager.layer_offsets = [7] + manager._window_size_by_layer_offset = {7: 128} + manager.impl = MagicMock() + manager.impl.get_cache_block_ids_range.return_value = [[20, 21]] + + result = manager.get_cache_indices_range(1, 4, 6, layer_idx=0) + + assert result == [20, 21] + manager.impl.get_cache_block_ids_range.assert_called_once_with(1, 128, 4, 6) + + def test_v1_manager_rejects_multiple_beams(self): + manager = object.__new__(KVCacheManager) + manager.layer_offsets = [0] + manager._window_size_by_layer_offset = {0: 128} + manager.impl = MagicMock() + manager.impl.get_cache_block_ids_range.return_value = [[20], [21]] + + with pytest.raises(ValueError, match="Chunked/pipelined KV transfer requires beam_width=1"): + manager.get_cache_indices_range(1, 0, 1, layer_idx=0) + + class TestPackedBeamBlockLayout: """Verify beam search block IDs stay 1-D with only final tail blocks appended.""" @@ -485,6 +620,9 @@ def _global_cached_token_count(self, req): # noqa: ARG002 def get_block_ids(self, req, group_idx, lg): # noqa: ARG002 return np.array([], dtype=np.int64) + def get_block_ids_range(self, req, group_idx, lg, block_begin, block_end): # noqa: ARG002 + return np.array([], dtype=np.int64) + def commit_blocks_for_reuse(self, req): # noqa: ARG002 pass diff --git a/tests/unittest/disaggregated/test_chunked_transfer.py b/tests/unittest/disaggregated/test_chunked_transfer.py index 1b9655549c1b..f2927d27ecb2 100644 --- a/tests/unittest/disaggregated/test_chunked_transfer.py +++ b/tests/unittest/disaggregated/test_chunked_transfer.py @@ -46,6 +46,7 @@ TxSession, WriteMeta, ) +from tensorrt_llm._torch.disaggregation.resource.page import AttentionLayerGroup from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequestState from tensorrt_llm.disaggregated_params import DisaggScheduleStyle from tensorrt_llm.llmapi.llm_args import CacheTransceiverConfig @@ -908,6 +909,9 @@ def test_pipelined_multiple_chunks_use_real_builder_and_tx_session(): transceiver._reuse_adapter = SimpleNamespace( tokens_per_block=tokens_per_block, get_block_ids=lambda _req, _idx, _lg: source_block_ids, + get_block_ids_range=( + lambda _req, _idx, _lg, block_begin, block_end: source_block_ids[block_begin:block_end] + ), ) transceiver._page_table = SimpleNamespace( layer_groups=[SimpleNamespace(sliding_window_size=None)] @@ -1080,22 +1084,28 @@ def _build_prefill_chunk_for( ): """Drive the real _build_prefill_chunk for one chunk of a prefilling request. - ``resident_blocks`` defaults to the chunk end, matching a source block list - that has only grown through the current chunk boundary. + The single full-attention layer group holds block ``i`` at ordinal ``i``. + ``resident_blocks`` bounds how far the cache manager has allocated and + defaults to the chunk end, matching a source list that has only grown + through the current chunk boundary. """ from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 if resident_blocks is None: resident_blocks = chunk_end_block - base_slice = KVSlice( - block_ids_per_layer_groups=[np.arange(resident_blocks, dtype=np.int64)], - ) transceiver = MagicMock() transceiver._kv_cache_manager.tokens_per_block = _REUSE_TPB - transceiver._create_kv_slice.return_value = base_slice + transceiver._page_table.layer_groups = [AttentionLayerGroup(pool_group_idx=0)] + transceiver._get_mamba_state_index.return_value = None transceiver._send_reqs = {} + def get_block_ids_range(_req, _group_idx, _layer_group, block_begin, block_end): + assert block_end <= resident_blocks, "callers must not read past the allocated blocks" + return np.arange(block_begin, block_end, dtype=np.int64) + + transceiver._reuse_adapter.get_block_ids_range.side_effect = get_block_ids_range + req = MagicMock() req.py_disaggregated_params = DisaggregatedParams(disagg_request_id=42) req.py_beam_width = 1 diff --git a/tests/unittest/disaggregated/test_kv_transfer.py b/tests/unittest/disaggregated/test_kv_transfer.py index 30d4396184b0..8d66e35dcf58 100644 --- a/tests/unittest/disaggregated/test_kv_transfer.py +++ b/tests/unittest/disaggregated/test_kv_transfer.py @@ -48,6 +48,7 @@ ) from tensorrt_llm._torch.disaggregation.native.transfer import TransferWorker, TransferWorkerConfig from tensorrt_llm._torch.disaggregation.resource.kv_extractor import KVRegionExtractorV1 +from tensorrt_llm._torch.disaggregation.resource.page import AttentionLayerGroup, MambaLayerGroup from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2 from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, LlmRequestType @@ -211,19 +212,44 @@ def _send_prefill_chunks( """ all_block_ids = [np.asarray(ids, dtype=np.int64) for ids in all_block_ids] total_blocks = max((len(ids) for ids in all_block_ids), default=0) - base_slice = KVSlice( - block_ids_per_layer_groups=all_block_ids, - mamba_state_index=mamba_state_index, - ) session = sender_session if sender_session is not None else MagicMock() session.kv_tasks = [] transceiver = MagicMock() transceiver._get_or_create_send_session.return_value = session - transceiver._create_kv_slice = MagicMock(return_value=base_slice) transceiver._reuse_adapter.tokens_per_block = tokens_per_block transceiver._kv_cache_manager.tokens_per_block = tokens_per_block transceiver._kv_cache_manager.kv_cache_map = {} transceiver._send_reqs = {} + # A layer group holding fewer blocks than the prompt is a sliding-window + # group whose leading blocks have been evicted, so its IDs are the resident + # *suffix* [total_blocks - len(ids), total_blocks) of the block table. + resident_begins = [total_blocks - len(block_ids) for block_ids in all_block_ids] + transceiver._page_table.layer_groups = [ + AttentionLayerGroup( + pool_group_idx=group_idx, + sliding_window_size=( + len(block_ids) * tokens_per_block if len(block_ids) < total_blocks else None + ), + ) + for group_idx, block_ids in enumerate(all_block_ids) + ] + if mamba_state_index is not None and transceiver._page_table.layer_groups: + transceiver._page_table.layer_groups[-1] = MambaLayerGroup( + pool_group_idx=len(all_block_ids) - 1 + ) + transceiver._get_mamba_state_index.return_value = mamba_state_index + + def get_block_ids_range(_req, group_idx, _layer_group, block_begin, block_end): + """Contiguous resident run ending at ``block_end`` — the adapter contract.""" + block_ids = all_block_ids[group_idx] + resident_begin = resident_begins[group_idx] + assert block_end <= total_blocks, "callers must not read past the allocated blocks" + begin = max(block_begin, resident_begin) + if begin >= block_end: + return block_ids[:0] + return block_ids[begin - resident_begin : block_end - resident_begin] + + transceiver._reuse_adapter.get_block_ids_range.side_effect = get_block_ids_range prompt_len = total_blocks * tokens_per_block req = MagicMock() @@ -276,6 +302,8 @@ def test_build_prefill_chunk_projects_incremental_source_against_full_prompt(): transceiver._kv_cache_manager.tokens_per_block = tokens_per_block transceiver._kv_cache_manager.kv_cache_map = {} transceiver._send_reqs = {} + transceiver._page_table.layer_groups = [AttentionLayerGroup(pool_group_idx=0)] + transceiver._get_mamba_state_index.return_value = None req = MagicMock() req.py_disaggregated_params = DisaggregatedParams(disagg_request_id=42) @@ -288,9 +316,7 @@ def test_build_prefill_chunk_projects_incremental_source_against_full_prompt(): chunk_start = chunk_idx * chunk_blocks chunk_end = chunk_start + chunk_blocks resident_ids = np.arange(chunk_end, dtype=np.int64) - transceiver._create_kv_slice.return_value = KVSlice( - block_ids_per_layer_groups=[resident_ids] - ) + transceiver._reuse_adapter.get_block_ids_range.return_value = resident_ids[-chunk_blocks:] req.py_last_context_chunk = ( chunk_start * tokens_per_block, chunk_end * tokens_per_block, @@ -305,27 +331,174 @@ def test_build_prefill_chunk_projects_incremental_source_against_full_prompt(): ) assert kv_slice.total_blocks == prompt_blocks + transceiver._create_kv_slice.assert_not_called() + assert transceiver._reuse_adapter.get_block_ids_range.call_count == 2 + assert all( + call.kwargs + == { + "block_begin": idx * chunk_blocks, + "block_end": (idx + 1) * chunk_blocks, + } + for idx, call in enumerate(transceiver._reuse_adapter.get_block_ids_range.call_args_list) + ) + + +def test_build_prefill_chunk_rejects_short_full_attention_range(): + transceiver = MagicMock() + transceiver._kv_cache_manager.tokens_per_block = 1 + transceiver._send_reqs = {} + transceiver._page_table.layer_groups = [AttentionLayerGroup(pool_group_idx=0)] + transceiver._reuse_adapter.get_block_ids_range.return_value = np.array([1], dtype=np.int64) + req = MagicMock( + py_beam_width=1, + prompt_len=4, + py_last_context_chunk=(0, 4), + context_remaining_length=0, + ) + req.py_disaggregated_params = DisaggregatedParams(disagg_request_id=42) + + with pytest.raises(ValueError, match="is not fully resident"): + KvCacheTransceiverV2._build_prefill_chunk(transceiver, req) + + +def test_build_prefill_chunk_rejects_overlong_range(): + """More blocks than the range spans means the group is not a run ending at chunk_end.""" + transceiver = MagicMock() + transceiver._kv_cache_manager.tokens_per_block = 1 + transceiver._send_reqs = {} + # A window wider than the prompt keeps the stale boundary at block 0, so the + # requested range is the whole chunk. + transceiver._page_table.layer_groups = [ + AttentionLayerGroup(pool_group_idx=0, sliding_window_size=8) + ] + transceiver._reuse_adapter.get_block_ids_range.return_value = np.array( + [1, 2, 3], dtype=np.int64 + ) + req = MagicMock( + py_beam_width=1, + prompt_len=4, + py_last_context_chunk=(0, 2), + context_remaining_length=2, + ) + req.py_disaggregated_params = DisaggregatedParams(disagg_request_id=42) + + with pytest.raises(ValueError, match="which spans only 2"): + KvCacheTransceiverV2._build_prefill_chunk(transceiver, req) + + +def test_build_prefill_chunk_accepts_short_swa_range(): + """A windowed group may return fewer blocks than its range spans.""" + transceiver = MagicMock() + transceiver._kv_cache_manager.tokens_per_block = 1 + transceiver._send_reqs = {} + transceiver._page_table.layer_groups = [ + AttentionLayerGroup(pool_group_idx=0, sliding_window_size=4) + ] + transceiver._reuse_adapter.get_block_ids_range.return_value = np.array([2, 3], dtype=np.int64) + transceiver._get_mamba_state_index.return_value = None + req = MagicMock( + py_beam_width=1, + prompt_len=4, + py_last_context_chunk=(0, 4), + context_remaining_length=0, + ) + req.py_disaggregated_params = DisaggregatedParams(disagg_request_id=42) + + kv_slice = KvCacheTransceiverV2._build_prefill_chunk(transceiver, req) + + # Range [1, 4): the window covers the prompt plus the first generated token. + assert transceiver._reuse_adapter.get_block_ids_range.call_args.kwargs == { + "block_begin": 1, + "block_end": 4, + } + np.testing.assert_array_equal(kv_slice.block_ids_per_layer_groups[0], [2, 3]) + + +def test_build_prefill_chunk_skips_windowed_group_before_final_window(): + """A chunk entirely below the final window contributes nothing for that group.""" + tokens_per_block = 4 + transceiver = MagicMock() + transceiver._kv_cache_manager.tokens_per_block = tokens_per_block + transceiver._send_reqs = {} + transceiver._page_table.layer_groups = [ + AttentionLayerGroup(pool_group_idx=0, sliding_window_size=3 * tokens_per_block) + ] + transceiver._get_mamba_state_index.return_value = None + + req = MagicMock( + py_beam_width=1, + prompt_len=16 * tokens_per_block, + prepopulated_prompt_len=0, + py_last_context_chunk=(4 * tokens_per_block, 8 * tokens_per_block), + context_remaining_length=8 * tokens_per_block, + ) + req.py_disaggregated_params = DisaggregatedParams(disagg_request_id=42) + + kv_slice = KvCacheTransceiverV2._build_prefill_chunk(transceiver, req) + + transceiver._reuse_adapter.get_block_ids_range.assert_not_called() + assert kv_slice.block_ids_per_layer_groups[0].size == 0 + + +def test_build_prefill_chunk_empty_range_skips_cache_queries(): + transceiver = MagicMock() + transceiver._kv_cache_manager.tokens_per_block = 1 + transceiver._send_reqs = {} + transceiver._page_table.layer_groups = [ + AttentionLayerGroup(pool_group_idx=0), + MambaLayerGroup(pool_group_idx=1), + ] + transceiver._get_mamba_state_index.return_value = 7 + req = MagicMock( + py_beam_width=1, + prompt_len=0, + py_last_context_chunk=(0, 0), + context_remaining_length=0, + ) + req.py_disaggregated_params = DisaggregatedParams(disagg_request_id=42) + + kv_slice = KvCacheTransceiverV2._build_prefill_chunk(transceiver, req) + + transceiver._reuse_adapter.get_block_ids_range.assert_not_called() + assert all(block_ids.size == 0 for block_ids in kv_slice.block_ids_per_layer_groups) + assert kv_slice.mamba_state_index == 7 + @pytest.mark.parametrize( - "source_block_ids", + "source_block_ids,resident_begin", [ - np.arange(16, dtype=np.int64), - np.arange(9, 13, dtype=np.int64), + (np.arange(16, dtype=np.int64), 0), + (np.arange(9, 13, dtype=np.int64), 9), ], ids=["v1_full_prompt_allocation", "v2_incremental_allocation"], ) -def test_build_prefill_chunk_normalizes_swa_source_to_computed_prefix(source_block_ids): - """A partial SWA chunk must not select full-prompt pages beyond its computed end.""" +def test_build_prefill_chunk_normalizes_swa_source_to_computed_prefix( + source_block_ids, resident_begin +): + """A partial SWA chunk must select only the block its final window shares with it. + + V1 reserves the whole prompt up front while V2 grows incrementally, so the two + hold different ordinal ranges; both must yield block 12 for chunk [11, 13). + """ tokens_per_block = 8 prompt_blocks = 16 window_blocks = 4 + requested_ranges = [] + + def get_block_ids_range(_req, _idx, _lg, block_begin, block_end): + """Resident run ending at ``block_end``, holding ordinal i at value i.""" + requested_ranges.append((block_begin, block_end)) + begin = max(block_begin, resident_begin) + if begin >= block_end: + return source_block_ids[:0] + return source_block_ids[begin - resident_begin : block_end - resident_begin] + layer_group = SimpleNamespace(sliding_window_size=window_blocks * tokens_per_block) transceiver = object.__new__(KvCacheTransceiverV2) transceiver._reuse_adapter = SimpleNamespace( tokens_per_block=tokens_per_block, - get_cached_token_count_per_layer_group=lambda req, layer_groups: [0], - get_block_ids=lambda req, idx, lg: source_block_ids, + get_block_ids_range=get_block_ids_range, ) transceiver._page_table = SimpleNamespace(layer_groups=[layer_group]) transceiver._kv_cache_manager = SimpleNamespace( @@ -346,6 +519,10 @@ def test_build_prefill_chunk_normalizes_swa_source_to_computed_prefix(source_blo kv_slice = KvCacheTransceiverV2._build_prefill_chunk(transceiver, req) + # The final window of a 16-block prompt plus its first generated token is + # [12, 16), so only block 12 is asked for -- reaching back to the chunk start + # would hand the sender worker a run it silently reinterprets as [12, 14). + assert requested_ranges == [(12, 13)] np.testing.assert_array_equal(kv_slice.block_ids_per_layer_groups[0], [12]) @@ -391,14 +568,19 @@ def test_send_prefill_chunks_integrity_check(prepopulated_blocks): def test_send_prefill_chunks_multiple_layer_groups(): - """Each source layer group is a resident suffix ending at the current chunk.""" + """A sliding-window group contributes only the part of its resident suffix in the chunk.""" all_block_ids = [list(range(8)), list(range(3))] slices = _send_prefill_chunks(all_block_ids, chunk_size_blocks=4) assert len(slices) == 2 assert np.array_equal(slices[0].block_ids_per_layer_groups[0], np.array([0, 1, 2, 3])) assert np.array_equal(slices[1].block_ids_per_layer_groups[0], np.array([4, 5, 6, 7])) - assert np.array_equal(slices[0].block_ids_per_layer_groups[1], np.array([0, 1, 2])) - assert np.array_equal(slices[1].block_ids_per_layer_groups[1], np.array([0, 1, 2])) + # The short group holds ordinals [5, 8), so chunk [0, 4) misses it entirely. + # Chunk [4, 8) carries only ordinals [6, 8): the final window also has to + # cover the first generated token, which puts its stale boundary one block + # past the currently resident head. Sending block 5 as well would make the + # sender worker pair it with destination block 6. + assert slices[0].block_ids_per_layer_groups[1].size == 0 + assert np.array_equal(slices[1].block_ids_per_layer_groups[1], np.array([1, 2])) assert slices[0].token_range == TokenRange(start=0, end=4) assert slices[1].token_range == TokenRange(start=4, end=8)