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/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..5103f520a5ab 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,30 @@ 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. + + 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: + 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 +98,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 +121,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 +161,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 +215,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/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 0d7e3429424b..aa6640fdc329 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -59,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) @@ -143,6 +146,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. @@ -251,10 +255,19 @@ def __exit__(self, _exc_type, _exc_val, _exc_tb): self.shutdown() 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. + """ adapter = self._reuse_adapter tpb = adapter.tokens_per_block assert self._page_table is not None layer_groups = self._page_table.layer_groups + resident_blocks = (req.prompt_len + tpb - 1) // tpb is_gen_only = req.is_generation_only_request() cached_per_lg = ( @@ -283,51 +296,49 @@ 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, ) 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.""" @@ -579,12 +590,35 @@ 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 - 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. + + 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.""" @@ -605,10 +639,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 +671,170 @@ 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 + 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}" + ) + + # 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) + 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( + 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=self._get_mamba_state_index(req), + 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 +873,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 +911,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 +956,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 +1138,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/_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/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_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 new file mode 100644 index 000000000000..f2927d27ecb2 --- /dev/null +++ b/tests/unittest/disaggregated/test_chunked_transfer.py @@ -0,0 +1,1516 @@ +# 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.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 + +# --------------------------------------------------------------------------- +# 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 + + 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() + + +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, + 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)] + ) + 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 +# --------------------------------------------------------------------------- + + +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_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. + + 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. + + 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 + + transceiver = MagicMock() + transceiver._kv_cache_manager.tokens_per_block = _REUSE_TPB + 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 + 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_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( + 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..8d66e35dcf58 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,8 @@ ) 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 from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager @@ -181,6 +185,423 @@ 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) + 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._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() + 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 = {} + 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) + 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._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, + ) + 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 + + 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,resident_begin", + [ + (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, 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_block_ids_range=get_block_ids_range, + ) + 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) + + # 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]) + + +@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(): + """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])) + # 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) + + +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 +1536,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 +1630,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 +1640,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 +1753,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 +1788,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 +2141,158 @@ 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, + token_range=TokenRange(start=0, end=request_len), + ) + 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"), + (2, 1, False, 2, 1, False, False, True, "v2_tp2_pp1_pipelined"), + (2, 1, False, 2, 1, False, False, False, "v1_tp2_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