diff --git a/tensorrt_llm/_torch/disaggregation/base/transfer.py b/tensorrt_llm/_torch/disaggregation/base/transfer.py index 320f7e067bdf..b7f2db9d3177 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, + total_blocks: 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, total_blocks)``. ``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, total_blocks - len(block_ids)) + resident_end = total_blocks + 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.""" @@ -65,6 +95,9 @@ 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 + chunk_block_offset: int = 0 + transfer_chunk_size: Optional[int] = None + total_blocks: Optional[int] = None class SessionStatus(Enum): @@ -158,7 +191,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. + The slice's ``chunk_block_offset`` field 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 f43dc49da160..79a82ff326b8 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -37,6 +37,7 @@ SessionStatus, TxSessionBase, WaitResult, + project_blocks_to_global_chunk, ) from tensorrt_llm._torch.disaggregation.native.auxiliary import AuxBuffer from tensorrt_llm._torch.disaggregation.native.messenger import ZMQMessenger, decode_message @@ -202,6 +203,16 @@ def __init__(self, params: DisaggregatedParams, slot: Optional[int]): class KVSendTask(SendTaskBase): + """A per-slice send task within a TxSession. + + Args: + kv_slice: The KV slice describing which blocks to transfer. + The slice's ``chunk_block_offset`` field indicates the + shared global chunk cursor in block units. + params: Disaggregated serving parameters for this request. + slice_id: Index of this slice within the session's task list. + """ + def __init__( self, kv_slice: KVSlice, @@ -209,7 +220,7 @@ def __init__( slice_id: int, prompt_len: Optional[int] = None, beam_width: int = 1, - ): + ) -> None: super().__init__(params) self.slice_id = slice_id self.transferred_count = 0 @@ -493,7 +504,7 @@ 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") ) @@ -503,7 +514,7 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): str(self._instance_rank).encode("ascii"), str(write_meta.unique_rid).encode("ascii"), str(write_meta.slice_id).encode("ascii"), - b"True", # is_last_slice — ensures receiver resolves its task future + b"True", # is_last_slice — ensures receiver resolves its task event AgentResult.FAILED.value.encode("ascii"), ] ) @@ -536,13 +547,19 @@ 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? + # The receiver always has a single monolithic task (slice_id=0). + # Sender-side chunking is transparent to the receiver: only the + # last chunk carries is_last_slice=True so the receiver knows + # when all data has arrived. Intermediate chunk results are + # sent (not suppressed) so that RDMA failures propagate to the + # receiver immediately rather than requiring a timeout. + receiver_slice_id = 0 self._get_or_connect_thread_dealer(write_meta.peer_endpoint).send( [ MessageType.KV_AGENT_RESULT, str(self._instance_rank).encode("ascii"), str(write_meta.unique_rid).encode("ascii"), - str(write_meta.slice_id).encode("ascii"), + str(receiver_slice_id).encode("ascii"), str(write_meta.is_last_slice).encode("ascii"), agent_result.value.encode("ascii"), ] @@ -701,10 +718,36 @@ def _build_kv_write_meta(self, task: KVSendTask, req_info: RecvReqInfo) -> Write dst_block_ids_per_groups = req_info.block_ids_per_layer_groups src_block_ids_per_groups = task._slice.block_ids_per_layer_groups - # Aggregate fragments from all matching pools using numpy concatenation + tpb = extractor.page_table.tokens_per_block + token_range = task._slice.token_range + slice_end = ( + task._prompt_len + if task._prompt_len is not None + else (token_range.end if token_range is not None else 0) + ) + total_blocks = task._slice.total_blocks + if total_blocks is None: + total_blocks = (slice_end + tpb - 1) // tpb + chunk_offset = task._slice.chunk_block_offset + chunk_block_count = task._slice.transfer_chunk_size + for (self_lg, self_pi), (peer_lg, peer_pi) in pool_mapping.items(): 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 chunk_offset > 0 or not task._slice.is_last_slice: + 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, + total_blocks=total_blocks, + ) + else: + dst_block_ids = full_dst_block_ids # Speculative decoding: generation may have one extra draft-token block. block_diff = dst_block_ids.size - src_block_ids.size @@ -719,15 +762,11 @@ def _build_kv_write_meta(self, task: KVSendTask, req_info: RecvReqInfo) -> Write f"src/dst block count mismatch: {src_block_ids.size} vs " f"{dst_block_ids.size} (expected diff <= 1)" ) - 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) # 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 src_beam0_blocks = Sender._beam0_block_count( src_block_ids, total_blocks, task._beam_width ) @@ -742,8 +781,16 @@ 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 + if chunk_block_count is not None and ( + chunk_offset > 0 or not task._slice.is_last_slice + ): + # Chunked lists are suffixes of the current global chunk, + # not of the full prompt. + src_start = (chunk_offset + chunk_block_count - src_beam0_blocks) * tpb + dst_start = (chunk_offset + chunk_block_count - dst_beam0_blocks) * tpb + else: + src_start = (total_blocks - src_beam0_blocks) * tpb + dst_start = (total_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: @@ -965,7 +1012,7 @@ def _respond_with_kv(self, _send_id: bytes, message: list[bytes]): 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; + # 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) @@ -1143,6 +1190,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 ) @@ -1760,15 +1811,15 @@ def process_kv_agent_result( ) 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); use the final KV task's + # expected transfer count so chunked sessions wait for all senders. with self.lock: if not self._kv_tasks: logger.warning( f"Aux result received before any KV tasks for request {self.request_id}" ) return - task = self._kv_tasks[0] + task = self._kv_tasks[-1] if status == AgentResult.SUCCESS: self._aux_count += 1 @@ -2024,7 +2075,18 @@ def populate_instance_and_rank_info(self, endpoints: list[str], layer_num_per_pp self._rank_info.sender_endpoints = endpoints self._rank_info.layer_num_per_pp = layer_num_per_pp - def create_tx_session(self, request: LlmRequest) -> TxSession: + def create_tx_session( + self, + request: LlmRequest, + ) -> TxSession: + """Create a TxSession for the given request. + + Args: + request: The LLM request to create a send session for. + + Returns: + A new ``TxSession`` ready to accept ``send()`` calls. + """ params = request.py_disaggregated_params assert params is not None return TxSession( diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index a5f11093e17b..674da392a0c4 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -1,3 +1,4 @@ +import math import time import uuid from collections import defaultdict @@ -16,6 +17,7 @@ TxSessionBase, WaitResult, get_unique_rid, + project_blocks_to_global_chunk, ) from tensorrt_llm._torch.disaggregation.native.transfer import TransferWorker, TransferWorkerConfig from tensorrt_llm._torch.disaggregation.resource.cache_reuse import ( @@ -98,6 +100,7 @@ def __init__( self._recv_reqs = {} self._wait_reqs = {} self._page_table = self._transfer_worker.page_table + self._transfer_chunk_size = cache_transceiver_config.transfer_chunk_size # Sticky role markers; flip True once any session opens, used to short-circuit # per-iter tp_allgather when this transceiver never sends/receives. @@ -239,6 +242,103 @@ def _create_kv_slice( token_range=token_range, ) + def _collect_base_slice(self, req: LlmRequest) -> KVSlice: + """Collect a full KVSlice (including metadata) for a request. + + This returns the complete slice produced by ``_create_kv_slice``, + preserving fields like ``mamba_state_index`` that are required + for hybrid-state model transfers. + + Args: + req: The LLM request whose KV cache block IDs to collect. + + Returns: + A ``KVSlice`` with all metadata populated. + """ + return self._create_kv_slice(req) + + def _create_kv_slices(self, req: LlmRequest) -> List[KVSlice]: + """Create one or more KVSlice objects for a request. + + When ``transfer_chunk_size`` is ``None``, returns a single slice + covering all blocks. Otherwise, each layer group's block ID + list is partitioned into slices of at most ``transfer_chunk_size`` + blocks. + + Args: + req: The LLM request to create slices for. + + Returns: + A list of ``KVSlice`` objects. Only the last slice has + ``is_last_slice=True``. + + Raises: + ValueError: If the reassembled block IDs from all slices do not + match the original block IDs. + """ + base_slice = self._collect_base_slice(req) + all_block_ids = base_slice.block_ids_per_layer_groups + + if self._transfer_chunk_size is None: + return [base_slice] + + max_resident_blocks = max((len(ids) for ids in all_block_ids), default=0) + if max_resident_blocks == 0: + return [base_slice] + + tpb = self._reuse_adapter.tokens_per_block + prompt_len = getattr(req, "prompt_len", None) + if isinstance(prompt_len, int) and prompt_len > 0: + total_blocks = math.ceil(prompt_len / tpb) + elif base_slice.token_range is not None: + total_blocks = math.ceil(base_slice.token_range.end / tpb) + else: + total_blocks = max_resident_blocks + total_blocks = max(max_resident_blocks, total_blocks) + + num_chunks = math.ceil(total_blocks / self._transfer_chunk_size) + slices: List[KVSlice] = [] + for chunk_idx in range(num_chunks): + start = chunk_idx * self._transfer_chunk_size + is_last = chunk_idx == num_chunks - 1 + chunk_block_count = min(self._transfer_chunk_size, total_blocks - start) + chunk_token_start = start * tpb + chunk_token_end = (start + chunk_block_count) * tpb + if base_slice.token_range is not None: + chunk_token_end = min(chunk_token_end, base_slice.token_range.end) + chunk_token_range = TokenRange(start=chunk_token_start, end=chunk_token_end) + + chunk_block_ids = [ + project_blocks_to_global_chunk( + ids, + chunk_block_offset=start, + chunk_block_count=chunk_block_count, + total_blocks=total_blocks, + ) + for ids in all_block_ids + ] + slices.append( + KVSlice( + is_last_slice=is_last, + block_ids_per_layer_groups=chunk_block_ids, + mamba_state_index=base_slice.mamba_state_index, + token_range=chunk_token_range, + chunk_block_offset=start, + transfer_chunk_size=chunk_block_count, + total_blocks=total_blocks, + ) + ) + + for lg_idx, original_ids in enumerate(all_block_ids): + reassembled = np.concatenate([s.block_ids_per_layer_groups[lg_idx] for s in slices]) + if not np.array_equal(reassembled, original_ids): + raise ValueError( + f"Chunking integrity check failed for layer group {lg_idx}: " + f"expected {len(original_ids)} blocks, got {len(reassembled)}" + ) + + return slices + @staticmethod def _split_packed_beam_block_ids( block_ids: np.ndarray, @@ -465,7 +565,8 @@ def respond_and_send_async(self, req: LlmRequest): self._ever_had_send_session = True session = self._get_or_create_send_session(req) req.state = LlmRequestState.DISAGG_CONTEXT_TRANS_IN_PROGRESS - session.send(self._create_kv_slice(req)) + for kv_slice in self._create_kv_slices(req): + session.send(kv_slice) self._finalize_send(req, session) @nvtx_range("KvCacheTransceiverV2.request_and_receive_sync") @@ -502,7 +603,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 rid = get_unique_rid(req) if rid in self._recv_sessions: @@ -513,7 +624,13 @@ def request_and_receive_async(self, req: LlmRequest): req.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS session = self._transfer_worker.create_rx_session(req) self._recv_sessions[rid] = session - session.receive(self._create_kv_slice(req)) + base_slice = self._collect_base_slice(req) + full_slice = KVSlice( + is_last_slice=True, + block_ids_per_layer_groups=base_slice.block_ids_per_layer_groups, + mamba_state_index=base_slice.mamba_state_index, + ) + session.receive(full_slice) self._recv_reqs[rid] = req def check_context_transfer_status( @@ -523,6 +640,7 @@ def check_context_transfer_status( self._ctx_need_tp_sync or self._ctx_need_pp_sync ): return [], [] + block_all = at_least_request_num is None wait_num = at_least_request_num if not block_all else 0 diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py index 8e6b47c8bff8..70bdede1a398 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py @@ -66,16 +66,63 @@ def create_kv_cache_transceiver( "MPI CacheTransceiver is deprecated, UCX or NIXL is recommended") elif cache_transceiver_config.backend == "UCX": logger.info( - f"Using UCX kv-cache transceiver. If your devices are not in the same domain, please consider setting " - f"UCX_CUDA_IPC_ENABLE_MNNVL=n, UCX_RNDV_SCHEME=put_zcopy and/or unset UCX_NET_DEVICES upon server " - f"hangs or lower-than-expected performance.") + "Using UCX kv-cache transceiver. If your devices are not in the same domain, please consider setting " + "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.") + + # Auto-select Python transceiver when transfer_chunk_size is set, + # since the C++ transceiver does not support chunked transfer. + # Only applies to NIXL/DEFAULT backends (the Python transceiver + # does not support UCX, MPI, or MOONCAKE). + runtime = cache_transceiver_config.transceiver_runtime + use_python = runtime == "PYTHON" + if (runtime is None + and cache_transceiver_config.transfer_chunk_size is not None): + if cache_transceiver_config.backend in (None, "DEFAULT", "NIXL"): + # Use warning (not info) so users notice the transceiver swap and + # the implied perf / staging-buffer characteristics change. Set + # transceiver_runtime='CPP' explicitly to opt out (and lose + # chunked transfer). + logger.warning( + "transfer_chunk_size is set; auto-selecting the Python " + "transceiver instead of the C++ transceiver to enable " + "chunked KV cache transfer. " + "Set transceiver_runtime='CPP' to disable this auto-selection.") + use_python = True + else: + logger.warning( + f"transfer_chunk_size is set but backend " + f"'{cache_transceiver_config.backend}' requires the C++ " + f"transceiver, which does not support chunked transfer. " + f"transfer_chunk_size will be ignored. Use NIXL backend to " + f"enable chunked transfer.") + elif (runtime == "CPP" + and cache_transceiver_config.transfer_chunk_size is not None): + raise ValueError( + "transfer_chunk_size is set but transceiver_runtime='CPP' " + "explicitly disables Python auto-selection; " + "transfer_chunk_size will be ignored.") + + # Warn when transfer_chunk_size is below the recommended floor. The Pydantic + # field is PositiveInt (>=1), but values below ~16 push the per-chunk RDMA + # overhead into the regime where it dominates transfer throughput. + _MIN_RECOMMENDED_TRANSFER_CHUNK_SIZE = 16 + if (cache_transceiver_config.transfer_chunk_size is not None + and cache_transceiver_config.transfer_chunk_size + < _MIN_RECOMMENDED_TRANSFER_CHUNK_SIZE): + logger.warning( + f"transfer_chunk_size={cache_transceiver_config.transfer_chunk_size} " + f"is below the recommended floor of " + f"{_MIN_RECOMMENDED_TRANSFER_CHUNK_SIZE}; per-chunk RDMA overhead " + f"may dominate transfer throughput. Consider 64-128 for " + f"long-context workloads (ISL >= 32K).") # Select transceiver implementation based on transceiver_runtime # transceiver_runtime == None or "CPP" -> use C++ transceiver (default) # transceiver_runtime == "PYTHON" -> use Python transceiver - if cache_transceiver_config.transceiver_runtime == "PYTHON": + if use_python: # Python transceiver currently only supports NIXL and DEFAULT backend - if cache_transceiver_config.backend not in ("DEFAULT", "NIXL"): + if cache_transceiver_config.backend not in (None, "DEFAULT", "NIXL"): raise ValueError( f"Python transceiver currently only supports NIXL or DEFAULT backend, " f"got {cache_transceiver_config.backend}. " diff --git a/tensorrt_llm/llmapi/llm_args.py b/tensorrt_llm/llmapi/llm_args.py index 3093f2f1e34f..03c43b68ca40 100644 --- a/tensorrt_llm/llmapi/llm_args.py +++ b/tensorrt_llm/llmapi/llm_args.py @@ -3539,7 +3539,23 @@ class CacheTransceiverConfig(StrictBaseModel, PybindMirror): "Bounded wait interval in milliseconds for polling KV transfer " "progress when active transfers block disaggregated admission.") + transfer_chunk_size: Optional[PositiveInt] = Field( + default=None, + description= + "Maximum number of KV cache blocks per layer group per chunk for " + "chunked KV cache transfer. When set, each layer group's block list " + "is partitioned into slices of at most this many blocks, and each " + "slice is transferred independently. The total data per chunk is " + "approximately transfer_chunk_size * num_layer_groups * slot_bytes. " + "This reduces per-transfer NIXL descriptor pressure for long " + "sequences. When None (default), the entire " + "KV cache is transferred in a single slice. When set with NIXL " + "backend (default), the Python transceiver is auto-selected. " + "Not supported with UCX, MPI, or MOONCAKE backends.") + def _to_pybind(self): + # transfer_chunk_size is consumed by the Python transceiver only + # and has no C++ counterpart, so it is intentionally omitted. return _CacheTransceiverConfig( backend=_CacheTransceiverBackendType.from_string(self.backend), max_tokens_in_buffer=self.max_tokens_in_buffer, diff --git a/tests/integration/defs/accuracy/test_disaggregated_serving.py b/tests/integration/defs/accuracy/test_disaggregated_serving.py index c1d1e370ea2a..56b3d92b1d3a 100644 --- a/tests/integration/defs/accuracy/test_disaggregated_serving.py +++ b/tests/integration/defs/accuracy/test_disaggregated_serving.py @@ -737,6 +737,49 @@ 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("transfer_chunk_size", [64]) + @parametrize_with_ids("enable_block_reuse", [False, True]) + def test_chunked_kv_transfer_nixl_python_accuracy(self, + transfer_chunk_size: int, + enable_block_reuse: bool): + """Test chunked 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, + "transfer_chunk_size": transfer_chunk_size, + } + ctx_server_config = { + "disable_overlap_scheduler": True, + "kv_cache_config": dict(kv_cache_config), + "cache_transceiver_config": dict(cache_transceiver_config), + } + gen_server_config = { + "disable_overlap_scheduler": False, + "kv_cache_config": dict(kv_cache_config), + "cache_transceiver_config": dict(cache_transceiver_config), + } + disaggregated_server_config = { + "hostname": "localhost", + "backend": "pytorch", + "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 = { 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 3c24a1809556..38d419da141d 100644 --- a/tests/integration/test_lists/test-db/l0_dgx_b200.yml +++ b/tests/integration/test_lists/test-db/l0_dgx_b200.yml @@ -17,6 +17,8 @@ l0_dgx_b200: tests: - 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::TestLlama3_1_8BInstruct::test_chunked_kv_transfer_nixl_python_accuracy[enable_block_reuse=False-transfer_chunk_size=64] + - accuracy/test_disaggregated_serving.py::TestLlama3_1_8BInstruct::test_chunked_kv_transfer_nixl_python_accuracy[enable_block_reuse=True-transfer_chunk_size=64] # ------------- 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/disaggregated/test_chunked_transfer.py b/tests/unittest/disaggregated/test_chunked_transfer.py new file mode 100644 index 000000000000..c997521d0e70 --- /dev/null +++ b/tests/unittest/disaggregated/test_chunked_transfer.py @@ -0,0 +1,390 @@ +# 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 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. +""" + +from types import SimpleNamespace +from unittest.mock import MagicMock + +import numpy as np + +from tensorrt_llm import DisaggregatedParams +from tensorrt_llm._torch.disaggregation.base.transfer import ( + KVSlice, + SessionStatus, + WaitResult, + project_blocks_to_global_chunk, +) +from tensorrt_llm._torch.disaggregation.native.transfer import ( + AgentResult, + KVSendTask, + RecvReqInfo, + RxSession, + Sender, + TaskStatus, + TxSession, +) + +# --------------------------------------------------------------------------- +# 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, **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(), + **kwargs, + ) + for i in range(num_slices): + s = KVSlice( + is_last_slice=(i == num_slices - 1), + block_ids_per_layer_groups=[[i]], + chunk_block_offset=i, + ) + session.send(s) + return session + + +def _make_rx_session(num_slices: int, rid: int = 42) -> 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(), + ) + 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 + + +# --------------------------------------------------------------------------- +# KVSendTask tests +# --------------------------------------------------------------------------- + + +def test_kv_send_task_chunk_block_offset(): + """KVSendTask reads chunk_block_offset from the slice.""" + s = KVSlice(is_last_slice=False, block_ids_per_layer_groups=[[0, 1]], chunk_block_offset=512) + task = KVSendTask(s, _make_params(), slice_id=1) + assert task._slice.chunk_block_offset == 512 + assert task.slice_id == 1 + assert task._slice is s + + +def test_kv_send_task_default_offset(): + """Default chunk_block_offset on KVSlice is 0.""" + s = KVSlice(is_last_slice=True, block_ids_per_layer_groups=[[0]]) + task = KVSendTask(s, _make_params(), slice_id=0) + assert task._slice.chunk_block_offset == 0 + + +# --------------------------------------------------------------------------- +# 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, + total_blocks=3, + ) + + assert projected_ids.size == 0 + + +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, + total_blocks=8, + ) + second_chunk = project_blocks_to_global_chunk( + block_ids, + chunk_block_offset=4, + chunk_block_count=4, + total_blocks=8, + ) + + assert first_chunk.size == 0 + assert np.array_equal(second_chunk, block_ids) + + +def test_build_kv_write_meta_projects_asymmetric_layer_group_chunk(): + """A short layer group's suffix blocks transfer with the overlapping global chunk.""" + 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 + + task = 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=SimpleNamespace(end=64), + chunk_block_offset=4, + transfer_chunk_size=4, + total_blocks=8, + ), + _make_params(), + slice_id=1, + ) + req_info = 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, + ) + + write_meta = sender._build_kv_write_meta(task, 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)) + + +# --------------------------------------------------------------------------- +# TxSession multi-slice status tests (real class) +# --------------------------------------------------------------------------- + + +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_uses_last_task(): + """process_aux_agent_result uses the last task's expected_transfers.""" + session = _make_rx_session(3) + session._kv_tasks[0].expected_transfers = 99 + session._kv_tasks[1].expected_transfers = 99 + session._kv_tasks[2].expected_transfers = 1 + + 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 diff --git a/tests/unittest/disaggregated/test_kv_transfer.py b/tests/unittest/disaggregated/test_kv_transfer.py index 130d074dd9fe..acf2df70e9c0 100644 --- a/tests/unittest/disaggregated/test_kv_transfer.py +++ b/tests/unittest/disaggregated/test_kv_transfer.py @@ -4,6 +4,7 @@ import random import time import uuid +from unittest.mock import MagicMock # Exclude IB (no fabric) and gdr_copy (UCX rcache SIGABRT at teardown). os.environ.setdefault("UCX_TLS", "^ib,gdr_copy") @@ -29,6 +30,7 @@ ) from tensorrt_llm._torch.disaggregation.native.transfer import TransferWorker, TransferWorkerConfig from tensorrt_llm._torch.disaggregation.resource.kv_extractor import KVRegionExtractorV1 +from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 from tensorrt_llm._torch.pyexecutor.kv_cache_manager_v2 import KVCacheManagerV2 from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, LlmRequestType from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager @@ -152,6 +154,95 @@ def test_session_status_enum(): assert len(SessionStatus) == 7 +# --------------------------------------------------------------------------- +# Chunked KV slice creation tests +# --------------------------------------------------------------------------- + + +def _chunk_block_ids(all_block_ids, transfer_chunk_size, mamba_state_index=None): + """Call the real _create_kv_slices via a mock transceiver.""" + from unittest.mock import MagicMock + + from tensorrt_llm._torch.disaggregation.base.transfer import KVSlice + + base_slice = KVSlice( + is_last_slice=True, + block_ids_per_layer_groups=all_block_ids, + mamba_state_index=mamba_state_index, + ) + + transceiver = MagicMock() + transceiver._transfer_chunk_size = transfer_chunk_size + transceiver._collect_base_slice = MagicMock(return_value=base_slice) + transceiver._reuse_adapter.tokens_per_block = 1 + transceiver._create_kv_slices = KvCacheTransceiverV2._create_kv_slices.__get__(transceiver) + + req = MagicMock() + return transceiver._create_kv_slices(req) + + +@pytest.mark.parametrize( + "all_block_ids,chunk_size,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_create_kv_slices_basic(all_block_ids, chunk_size, expected_num_slices): + """Chunking produces the expected number of slices.""" + slices = _chunk_block_ids(all_block_ids, transfer_chunk_size=chunk_size) + 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 + + +def test_create_kv_slices_integrity_check(): + """Reassembled block IDs from all slices must match the original.""" + all_block_ids = [list(range(17)), list(range(5))] + slices = _chunk_block_ids(all_block_ids, transfer_chunk_size=4) + 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_create_kv_slices_multiple_layer_groups(): + """Shorter layer groups are projected into the overlapping global chunk.""" + all_block_ids = [list(range(8)), list(range(3))] + slices = _chunk_block_ids(all_block_ids, transfer_chunk_size=4) + assert len(slices) == 2 + assert np.array_equal(slices[0].block_ids_per_layer_groups[0], np.array([0, 1, 2, 3])) + assert np.array_equal(slices[1].block_ids_per_layer_groups[0], np.array([4, 5, 6, 7])) + assert len(slices[0].block_ids_per_layer_groups[1]) == 0 + assert np.array_equal(slices[1].block_ids_per_layer_groups[1], np.array([0, 1, 2])) + assert slices[0].token_range == TokenRange(start=0, end=4) + assert slices[1].token_range == TokenRange(start=4, end=8) + + +def test_create_kv_slices_preserves_mamba_state_index(): + """mamba_state_index is propagated to every chunk slice.""" + all_block_ids = [list(range(8))] + slices = _chunk_block_ids(all_block_ids, transfer_chunk_size=4, mamba_state_index=42) + assert len(slices) == 2 + for s in slices: + assert s.mamba_state_index == 42 + + +def test_create_kv_slices_none_mamba_state_index(): + """mamba_state_index=None is preserved when not set.""" + all_block_ids = [list(range(4))] + slices = _chunk_block_ids(all_block_ids, transfer_chunk_size=4) + assert len(slices) == 1 + assert slices[0].mamba_state_index is None + + def create_transfer_worker_setup( ctx_tp: int, ctx_pp: int, @@ -1078,7 +1169,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( + "transfer_chunk_size", + [None, 2], + ids=["single_slice", "sender_chunked"], +) +def test_transfer_with_gen_prefix_offset(use_v2, transfer_chunk_size): """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), @@ -1167,13 +1263,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) @@ -1183,7 +1273,31 @@ 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 transfer_chunk_size 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: + transceiver = MagicMock() + transceiver._transfer_chunk_size = transfer_chunk_size + transceiver._collect_base_slice = MagicMock( + return_value=KVSlice( + is_last_slice=True, + block_ids_per_layer_groups=ctx_block_ids, + token_range=TokenRange(start=0, end=request_len), + ) + ) + transceiver._reuse_adapter.tokens_per_block = tokens_per_block + transceiver._create_kv_slices = KvCacheTransceiverV2._create_kv_slices.__get__( + transceiver + ) + for kv_slice in transceiver._create_kv_slices(ctx_request): + tx.send(kv_slice) result = tx.wait_complete() assert result == WaitResult.COMPLETED, f"tx wait_complete returned {result}" @@ -1281,7 +1395,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, @@ -1316,19 +1430,236 @@ 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, + transfer_chunk_size, +): + """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): + transceiver = MagicMock() + transceiver._transfer_chunk_size = transfer_chunk_size + transceiver._collect_base_slice = MagicMock( + return_value=KVSlice( + is_last_slice=True, + block_ids_per_layer_groups=block_ids_per_groups, + token_range=token_range, + ) + ) + transceiver._reuse_adapter.tokens_per_block = setup["tokens_per_block"] + transceiver._create_kv_slices = KvCacheTransceiverV2._create_kv_slices.__get__(transceiver) + for kv_slice in transceiver._create_kv_slices(ctx_info["ctx_request"]): + sender_session.send(kv_slice) + + 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, transfer_chunk_size=chunk_size) + add_and_verify_chunked_request(setup, 2, 3, request_len * 2, transfer_chunk_size=chunk_size) + finally: + for worker in setup["ctx_transfer_workers"]: + worker.shutdown() for worker in setup["gen_transfer_workers"]: worker.shutdown() diff --git a/tests/unittest/llmapi/test_llm_args.py b/tests/unittest/llmapi/test_llm_args.py index 69bab9c13e00..7cf6b132bcfb 100644 --- a/tests/unittest/llmapi/test_llm_args.py +++ b/tests/unittest/llmapi/test_llm_args.py @@ -1827,6 +1827,23 @@ def test_cache_transceiver_config_arbitrary_args(self): CacheTransceiverConfig(backend="UCX", invalid_config="should_fail") assert "invalid_config" in str(exc_info.value) + def test_cache_transceiver_config_transfer_chunk_size(self): + """Test transfer_chunk_size field validation.""" + config = CacheTransceiverConfig(transfer_chunk_size=64) + assert config.transfer_chunk_size == 64 + + config_none = CacheTransceiverConfig(transfer_chunk_size=None) + assert config_none.transfer_chunk_size is None + + config_default = CacheTransceiverConfig() + assert config_default.transfer_chunk_size is None + + with pytest.raises(pydantic_core._pydantic_core.ValidationError): + CacheTransceiverConfig(transfer_chunk_size=0) + + with pytest.raises(pydantic_core._pydantic_core.ValidationError): + CacheTransceiverConfig(transfer_chunk_size=-1) + def test_torch_compile_config_arbitrary_args(self): """Test that TorchCompileConfig rejects arbitrary arguments.""" # Valid arguments should work