From ec0875a42f56033637e832b890520eed2e3efe09 Mon Sep 17 00:00:00 2001 From: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> Date: Fri, 20 Mar 2026 03:19:05 +0000 Subject: [PATCH 1/5] update the worker and transceiver Signed-off-by: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> --- .../_torch/disaggregation/base/agent.py | 5 +- .../native/py_cache_transceiver.py | 434 ------------------ .../_torch/disaggregation/native/rank_info.py | 42 ++ .../_torch/disaggregation/native/transfer.py | 240 ++++------ .../_torch/disaggregation/resource/utils.py | 24 + .../_torch/disaggregation/transceiver.py | 386 ++++++++++++++++ .../_torch/pyexecutor/kv_cache_transceiver.py | 11 +- .../disaggregated/test_kv_transfer.py | 50 +- .../disaggregated/test_kv_transfer_mp.py | 30 +- .../test_py_cache_transceiver_mp.py | 25 +- 10 files changed, 589 insertions(+), 658 deletions(-) delete mode 100644 tensorrt_llm/_torch/disaggregation/native/py_cache_transceiver.py create mode 100644 tensorrt_llm/_torch/disaggregation/transceiver.py diff --git a/tensorrt_llm/_torch/disaggregation/base/agent.py b/tensorrt_llm/_torch/disaggregation/base/agent.py index 1ac8647377fd..bf4899707e03 100644 --- a/tensorrt_llm/_torch/disaggregation/base/agent.py +++ b/tensorrt_llm/_torch/disaggregation/base/agent.py @@ -1,7 +1,7 @@ import os from abc import ABC, abstractmethod from dataclasses import dataclass -from typing import List, NamedTuple, Optional +from typing import List, NamedTuple, Optional, Tuple from tensorrt_llm import logger @@ -25,7 +25,6 @@ class MemoryDesc(NamedTuple): ptr: int size: int device_id: int - name: Optional[str] = None @dataclass @@ -46,7 +45,7 @@ class TransferRequest: @dataclass class RegMemoryDescs: type: str - descs: List[MemoryDesc] + descs: List[Tuple[int, int, int, str]] class TransferStatus(ABC): diff --git a/tensorrt_llm/_torch/disaggregation/native/py_cache_transceiver.py b/tensorrt_llm/_torch/disaggregation/native/py_cache_transceiver.py deleted file mode 100644 index 4b2570d83bd0..000000000000 --- a/tensorrt_llm/_torch/disaggregation/native/py_cache_transceiver.py +++ /dev/null @@ -1,434 +0,0 @@ -import uuid -from collections import defaultdict -from itertools import chain -from typing import Any, Dict, List - -import torch - -import tensorrt_llm -from tensorrt_llm import logger -from tensorrt_llm._torch.disaggregation.base.transfer import KVSlice, WaitResult, get_unique_rid -from tensorrt_llm._torch.disaggregation.native.auxiliary import AuxBuffer -from tensorrt_llm._torch.disaggregation.native.transfer import TransferWorker -from tensorrt_llm._torch.disaggregation.resource.utils import get_global_layer_ids -from tensorrt_llm._torch.distributed.communicator import Distributed -from tensorrt_llm._torch.pyexecutor.kv_cache_transceiver import KvCacheTransceiver -from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest -from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager -from tensorrt_llm.bindings import LlmRequestState -from tensorrt_llm.bindings.executor import ContextPhaseParams -from tensorrt_llm.disaggregated_params import DisaggScheduleStyle -from tensorrt_llm.llmapi.llm_args import CacheTransceiverConfig -from tensorrt_llm.mapping import Mapping - -CacheTransceiverCpp = tensorrt_llm.bindings.internal.batch_manager.CacheTransceiver -AttentionTypeCpp = tensorrt_llm.bindings.internal.batch_manager.AttentionType -CacheTransBufferManagerCpp = tensorrt_llm.bindings.internal.batch_manager.CacheTransBufferManager -BackendTypeCpp = tensorrt_llm.bindings.executor.CacheTransceiverBackendType - - -def _find_consensus_request_ids(request_ids_all_ranks, sync_size): - frequency_map = defaultdict(int) - consensus_request_ids = [] - for request_id in list(chain.from_iterable(request_ids_all_ranks)): - frequency_map[request_id] += 1 - sorted_frequency_map = sorted(frequency_map.items(), key=lambda x: x[1], reverse=True) - for request_id, frequency in sorted_frequency_map: - if frequency == sync_size: - consensus_request_ids.append(request_id) - else: - break - return consensus_request_ids - - -class PyNativeCacheTransceiver(KvCacheTransceiver): - def __init__( - self, - mapping: Mapping, - dist: Distributed, - kv_cache_manager: KVCacheManager, - attention_type: AttentionTypeCpp, - cache_transceiver_config: CacheTransceiverConfig, - ): - self.dist: Distributed = dist - self.kv_cache_manager = kv_cache_manager - self.kv_transfer_timeout_ms = cache_transceiver_config.kv_transfer_timeout_ms - self.mapping = mapping - self._check_compatible() - self.sender_future_timeout_ms = ( - cache_transceiver_config.kv_transfer_sender_future_timeout_ms - ) - instance_name = None - if dist.rank == 0: - instance_name = str(uuid.uuid4()) - dist.broadcast(instance_name, 0) - else: - instance_name = dist.broadcast(instance_name, 0) - - self.instance_name = instance_name - - # device_id = mapping.local_rank - self.device_id = torch.cuda.current_device() - logger.info(f"device_id: {self.device_id} in PyNativeCacheTransceiver") - - # Aux payload carries first-gen and draft tokens in generation-first flow. - self.aux_buffer = AuxBuffer( - # * 2 to allow back-to-back batches, one in transferring, one in preparing next batch - max_slot_num=max(1, int(self.kv_cache_manager.max_batch_size)) * 2, - beam_width=max(1, int(getattr(self.kv_cache_manager, "max_beam_width", 1))), - max_draft_len=max(0, int(getattr(self.kv_cache_manager, "max_draft_len", 0))), - device="cpu", - ) - - self.transfer_worker = TransferWorker( - kv_cache_manager=kv_cache_manager, - mapping=mapping, - device_id=self.device_id, - instance_name=instance_name, - aux_buffer=self.aux_buffer, - ) - - self.context_info_endpoint = None - self.dp_rank = self.mapping.tp_rank if self.mapping.enable_attention_dp else 0 - if self.dist.rank == 0: - self.context_info_endpoint = self.transfer_worker.rank_info_server_endpoint - self.dist.broadcast(self.context_info_endpoint, 0) - else: - self.context_info_endpoint = self.dist.broadcast(self.context_info_endpoint, 0) - - self.mapping = mapping - - self.ctx_need_tp_sync = mapping.tp_size > 1 and (not mapping.enable_attention_dp) - - self.gen_need_sync = not ( - mapping.world_size == 1 or (mapping.enable_attention_dp and mapping.pp_size == 1) - ) - self.gen_sync_allgather_fun = ( - self.dist.pp_allgather if mapping.enable_attention_dp else self.dist.allgather - ) - ctx_server_endpoint = self.transfer_worker.sender_endpoint - layer_num = len(self.kv_cache_manager.pp_layers) - - ctx_server_endpoints = self.dist.allgather(ctx_server_endpoint) - layer_num_per_pp = self.dist.pp_allgather(layer_num) - self.transfer_worker.populate_instance_and_rank_info( - endpoints=ctx_server_endpoints, layer_num_per_pp=layer_num_per_pp - ) - - logger.info(f"transfer worker ctx_server_endpoints: {ctx_server_endpoints}") - logger.info(f"layer_num_per_pp: {layer_num_per_pp}") - logger.info(f"self.context_info_endpoint: {self.context_info_endpoint}") - self.send_sessions = {} # request_id to send_session - self.recv_sessions = {} # request_id to recv_session - self.send_req_id_to_request = {} # request_id to request (for send) - self.recv_req_id_to_request = {} # request_id to request (for recv) - self.wait_req_id_to_request = {} # request_id to request (for gen-first waiting-scheduler) - self.page_table = self.transfer_worker.page_table - # Check if using V2 manager (has kv_cache_map attribute) - self.is_v2_manager = hasattr(self.kv_cache_manager, "kv_cache_map") - - def shutdown(self): - if self.transfer_worker is not None: - self.transfer_worker.shutdown() - - def _create_kv_slice(self, req: LlmRequest): - # Get block_ids for each layer group - block_ids_per_layer_groups: List[List[int]] = [] - tokens_per_block = self.kv_cache_manager.tokens_per_block - - for group_idx, lg in enumerate(self.page_table.layer_groups): - if self.is_v2_manager: - # V2: Use get_aggregated_page_indices for efficient slot indices - group_id = group_idx - block_ids = list( - self.kv_cache_manager.kv_cache_map[ - req.py_request_id - ].get_aggregated_page_indices(group_id, valid_only=True) - ) - else: - # V1: Use get_batch_cache_indices - first_global_layer_id = get_global_layer_ids(lg)[0] - block_ids = self.kv_cache_manager.get_batch_cache_indices( - [req.py_request_id], layer_idx=first_global_layer_id - )[0] - - # Filter to only window-relevant blocks for sliding window layer groups. - # Computes the expected number of non-stale blocks (using the same - # eviction formula as update_resources) and keeps only the tail. - # This works correctly regardless of whether update_resources has - # been called: - # - Pre-eviction: all blocks present → trim to last N. - # - Post-eviction (V2 valid_only=True): stale blocks already - # removed → len == expected_valid, so the condition is false. - window_size = lg.sliding_window_size - if window_size is not None: - total_blocks = (req.prompt_len + tokens_per_block - 1) // tokens_per_block - stale_end = max(0, (req.prompt_len + 1 - window_size) // tokens_per_block) - expected_valid = total_blocks - stale_end - if expected_valid <= 0: - block_ids = [] - elif len(block_ids) > expected_valid: - block_ids = block_ids[-expected_valid:] - - block_ids_per_layer_groups.append(list(block_ids)) - - return KVSlice(is_last_slice=True, block_ids_per_layer_groups=block_ids_per_layer_groups) - - @staticmethod - def _need_aux_transfer(req: LlmRequest) -> bool: - params = req.py_disaggregated_params - return params is not None and params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST - - # starts background transfer to send this request's KV cache to the gen server, attaches ContextPhaseParams metadata - def respond_and_send_async(self, req: LlmRequest): - unique_rid = get_unique_rid(req) - if unique_rid not in self.send_sessions: - send_session = self.transfer_worker.create_tx_session(req) - self.send_sessions[unique_rid] = send_session - else: - send_session = self.send_sessions[unique_rid] - req.state = LlmRequestState.DISAGG_CONTEXT_TRANS_IN_PROGRESS - - # stores block ids, not raw data - kv_slice = self._create_kv_slice(req) - # sending actual kv data - send_session.send(kv_slice) - if self._need_aux_transfer(req): - send_session.pack_aux(req) - send_session.send_aux() - - # contains metadata about itself so the gen server can see - req.context_phase_params = ContextPhaseParams( - first_gen_tokens=[], - req_id=unique_rid, - opaque_state=None, - draft_tokens=None, - ctx_dp_rank=self.dp_rank, - disagg_info_endpoint=self.context_info_endpoint, - ) - self.send_req_id_to_request[unique_rid] = req - - return - - def request_and_receive_sync(self, req: LlmRequest): - raise NotImplementedError("request_and_receive_sync is not implemented") - - # starts background listener to receive KV cache from the ctx server into this request's pre-allocated blocks. - def request_and_receive_async(self, req: LlmRequest): - unique_rid = get_unique_rid(req) - req.state = LlmRequestState.DISAGG_GENERATION_TRANS_IN_PROGRESS - - # create rx session for receiving blocks - recv_session = self.transfer_worker.create_rx_session(req) - self.recv_sessions[unique_rid] = recv_session - kv_slice = self._create_kv_slice(req) - recv_session.receive(kv_slice) - self.recv_req_id_to_request[unique_rid] = req - - def check_context_transfer_status(self, at_least_request_num: int, mark_complete: bool = False): - block_all = at_least_request_num is None - - wait_num = at_least_request_num if not block_all else 0 - - local_completed_request_ids = [] - local_failed_request_ids = [] - for request_id, session in self.send_sessions.items(): - req = self.send_req_id_to_request[request_id] - need_aux = self._need_aux_transfer(req) - if session.is_completed(need_aux): - local_completed_request_ids.append(request_id) - elif session.has_failed(): - local_failed_request_ids.append(request_id) - local_sync_request_ids = local_completed_request_ids + local_failed_request_ids - - if self.ctx_need_tp_sync: - sync_request_ids_all_ranks = self.dist.tp_allgather(local_sync_request_ids) - else: - sync_request_ids_all_ranks = [local_sync_request_ids] - - sync_size = self.dist.tp_size if self.ctx_need_tp_sync else 1 - - to_complete_request_ids = _find_consensus_request_ids(sync_request_ids_all_ranks, sync_size) - for request_id in self.send_req_id_to_request.keys(): - if len(to_complete_request_ids) >= wait_num: - break - if request_id not in to_complete_request_ids: - to_complete_request_ids.append(request_id) - if block_all: - to_complete_request_ids = self.send_req_id_to_request.keys() - completed_request_ids = [] - timeout_request_ids = [] - failed_request_ids = [] - timeout = self.sender_future_timeout_ms / 1000.0 - for request_id in to_complete_request_ids: - session = self.send_sessions[request_id] - req = self.send_req_id_to_request[request_id] - result = session.wait_complete(need_aux=self._need_aux_transfer(req), timeout=timeout) - if result == WaitResult.COMPLETED: - completed_request_ids.append(request_id) - elif result == WaitResult.TIMEOUT: - logger.warning( - f"TxSession {session.disagg_request_id} timed out waiting for completion" - ) - timeout_request_ids.append(request_id) - else: - logger.warning(f"TxSession {session.disagg_request_id} failed to complete") - failed_request_ids.append(request_id) - - for request_id in completed_request_ids + failed_request_ids: - if request_id in completed_request_ids: - if mark_complete: - self.send_req_id_to_request[ - request_id - ].state = LlmRequestState.DISAGG_CONTEXT_COMPLETE - elif request_id in failed_request_ids: - self.send_req_id_to_request[request_id].state = LlmRequestState.DISAGG_TRANS_ERROR - del self.send_req_id_to_request[request_id] - self.transfer_worker.clear_session(self.send_sessions[request_id]) - del self.send_sessions[request_id] - - return completed_request_ids, failed_request_ids - - def check_gen_transfer_status(self, at_least_request_num: int): - block_all = at_least_request_num is None - - wait_num = at_least_request_num if not block_all else 0 - - local_completed_request_ids = [] - local_failed_request_ids = [] - for request_id, session in self.recv_sessions.items(): - req = self.recv_req_id_to_request[request_id] - need_aux = self._need_aux_transfer(req) - if session.is_completed(need_aux): - local_completed_request_ids.append(request_id) - elif session.has_failed(): - local_failed_request_ids.append(request_id) - local_sync_request_ids = local_completed_request_ids + local_failed_request_ids - - if self.gen_need_sync: - sync_request_ids = self.gen_sync_allgather_fun(local_sync_request_ids) - else: - sync_request_ids = [local_sync_request_ids] - - frequency_map = {} - for request_id in list(chain.from_iterable(sync_request_ids)): - frequency_map[request_id] = frequency_map.get(request_id, 0) + 1 - sorted_frequency_map = sorted(frequency_map.items(), key=lambda x: x[1], reverse=True) - sync_size = ( - self.mapping.pp_size if self.mapping.enable_attention_dp else self.mapping.world_size - ) - to_complete_request_ids = [] - for request_id, frequency in sorted_frequency_map: - if frequency == sync_size: - to_complete_request_ids.append(request_id) - else: - break - for request_id in self.recv_sessions.keys(): - if len(to_complete_request_ids) >= wait_num: - break - if request_id not in to_complete_request_ids: - to_complete_request_ids.append(request_id) - if block_all: - to_complete_request_ids = list(self.recv_sessions.keys()) - completed_request_ids = [] - failed_request_ids = [] - for request_id in to_complete_request_ids: - recv_session = self.recv_sessions[request_id] - req = self.recv_req_id_to_request[request_id] - result = recv_session.wait_complete( - need_aux=self._need_aux_transfer(req), block_for_aux=block_all - ) - if result == WaitResult.COMPLETED: - completed_request_ids.append(request_id) - elif result == WaitResult.FAILED: - failed_request_ids.append(request_id) - # else: None — KV done but aux still in flight; re-poll next cycle - - for request_id in completed_request_ids + failed_request_ids: - recv_session = self.recv_sessions[request_id] - req = self.recv_req_id_to_request[request_id] - if request_id in completed_request_ids: - if self._need_aux_transfer(req): - recv_session.unpack_aux(req) - first_gen_tokens = req.py_first_gen_tokens - draft_tokens = req.py_draft_tokens - if req.context_phase_params is None: - req.context_phase_params = ContextPhaseParams( - first_gen_tokens=first_gen_tokens, - req_id=req.py_request_id, - opaque_state=b"", - draft_tokens=draft_tokens, - ctx_dp_rank=0, - disagg_info_endpoint="", - ) - else: - req.context_phase_params.first_gen_tokens = first_gen_tokens - req.context_phase_params.draft_tokens = draft_tokens - req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE - elif request_id in failed_request_ids: - req.state = LlmRequestState.DISAGG_TRANS_ERROR - del self.recv_req_id_to_request[request_id] - self.transfer_worker.clear_session(recv_session) - del self.recv_sessions[request_id] - - return - - def check_gen_transfer_complete(self): - return len(self.recv_sessions) == 0 - - def cancel_request(self, req: LlmRequest): - raise NotImplementedError("cancel_request is not implemented") - - def get_disaggregated_params(self) -> Dict[str, Any]: - # Keep this aligned with fields populated in respond_and_send_async(). - # These values are server-level metadata used to seed generation-first - # requests before context-phase response data arrives. - return { - "ctx_dp_rank": self.dp_rank, - "ctx_info_endpoint": [self.context_info_endpoint] - if self.context_info_endpoint - else None, - } - - def prepare_context_requests(self, requests: List[LlmRequest]): - # Place new generation-first context requests into wait state, then - # use tp_allgather consensus to promote ready requests to CONTEXT_INIT. - for req in requests: - unique_rid = get_unique_rid(req) - if unique_rid not in self.send_sessions: - self.wait_req_id_to_request[unique_rid] = req - req.state = LlmRequestState.DISAGG_CONTEXT_WAIT_SCHEDULER - - # Check which waiting requests have peer info locally, then use - # tp_allgather consensus so all TP ranks agree before promoting. - # Without consensus, background peer info arriving at different - # times on different ranks causes scheduling mismatches → hang. - # Place tp sync here because this function runs in every iteration - # but check_context_transfer_status runs when can_queue is True - local_ready_request_ids = [] - for request_id in self.wait_req_id_to_request.keys(): - if self.transfer_worker.has_all_peer_req_infos_for_send(request_id): - local_ready_request_ids.append(request_id) - - if self.ctx_need_tp_sync: - ready_request_ids_all_ranks = self.dist.tp_allgather(local_ready_request_ids) - else: - ready_request_ids_all_ranks = [local_ready_request_ids] - - sync_size = self.dist.tp_size if self.ctx_need_tp_sync else 1 - ready_request_ids = _find_consensus_request_ids(ready_request_ids_all_ranks, sync_size) - - for request_id in ready_request_ids: - self.wait_req_id_to_request[request_id].state = LlmRequestState.CONTEXT_INIT - del self.wait_req_id_to_request[request_id] - - def _check_compatible(self): - if self.mapping.cp_size != 1: - raise ValueError( - f"PyNativeCacheTransceiver: _check_compatible: only support context parallelism is 1: " - f"cp_size: {self.mapping.cp_size}" - ) - return - - def get_context_state(self): - raise NotImplementedError("get_context_state is not implemented") diff --git a/tensorrt_llm/_torch/disaggregation/native/rank_info.py b/tensorrt_llm/_torch/disaggregation/native/rank_info.py index e6e4fa64ccd5..c908b7e415da 100644 --- a/tensorrt_llm/_torch/disaggregation/native/rank_info.py +++ b/tensorrt_llm/_torch/disaggregation/native/rank_info.py @@ -5,7 +5,9 @@ from tensorrt_llm._torch.disaggregation.native.auxiliary import AuxBufferMeta from tensorrt_llm._torch.disaggregation.native.mixers.attention.spec import AttentionInfo +from tensorrt_llm._torch.disaggregation.resource.kv_extractor import build_page_table_from_manager from tensorrt_llm._torch.disaggregation.resource.page import KVCachePageTable +from tensorrt_llm._utils import get_size_in_bytes @dataclass @@ -45,6 +47,46 @@ def to_bytes(self) -> bytes: data["page_table"] = self.page_table.to_dict() if self.page_table is not None else None return msgpack.packb(data) + @classmethod + def from_kv_cache_manager( + cls, + instance_name: str, + kv_cache_manager, + device_id: int, + aux_buffer_meta: Optional[AuxBufferMeta] = None, + ) -> "RankInfo": + m = kv_cache_manager.mapping + kvm = kv_cache_manager + enable_attention_dp = m.enable_attention_dp + return cls( + instance_name=instance_name, + instance_rank=m.rank, + tp_size=m.tp_size, + tp_rank=m.tp_rank, + pp_size=m.pp_size, + pp_rank=m.pp_rank, + dp_size=m.tp_size if enable_attention_dp else m.dp_size, + dp_rank=m.tp_rank if enable_attention_dp else 0, + cp_size=m.cp_size, + cp_rank=m.cp_rank, + device_id=device_id, + layer_num_per_pp=[len(kvm.pp_layers)], + sender_endpoints=[], + server_endpoint="", + self_endpoint="", + transfer_engine_info=bytes(), + attention=AttentionInfo( + kv_heads_per_rank=kvm.num_kv_heads_per_layer[0], + tokens_per_block=kvm.tokens_per_block, + dims_per_head=kvm.head_dim, + element_bytes=get_size_in_bytes(1, kvm.dtype), + enable_attention_dp=enable_attention_dp, + is_mla=kvm.kv_factor == 1, + ), + aux_meta=aux_buffer_meta, + page_table=build_page_table_from_manager(kvm), + ) + @classmethod def from_bytes(cls, data: bytes) -> "RankInfo": unpacked = msgpack.unpackb(data, strict_map_key=False) diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index df6234641193..dfca6dab0e9c 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -19,9 +19,10 @@ from cuda import cudart import tensorrt_llm.bindings -from tensorrt_llm import Mapping, logger +from tensorrt_llm import logger from tensorrt_llm._torch.disaggregation.base.agent import ( BaseTransferAgent, + MemoryDesc, MemoryDescs, MemoryType, RegMemoryDescs, @@ -40,20 +41,16 @@ ) from tensorrt_llm._torch.disaggregation.native.auxiliary import AuxBuffer from tensorrt_llm._torch.disaggregation.native.messenger import ZMQMessenger, decode_message -from tensorrt_llm._torch.disaggregation.native.mixers.attention.spec import AttentionInfo from tensorrt_llm._torch.disaggregation.native.peer import PeerRegistrar from tensorrt_llm._torch.disaggregation.native.perf_logger import PerfTimer, perf_log_manager from tensorrt_llm._torch.disaggregation.native.rank_info import RankInfo from tensorrt_llm._torch.disaggregation.native.utils import get_local_ip from tensorrt_llm._torch.disaggregation.nixl.agent import NixlTransferAgent -from tensorrt_llm._torch.disaggregation.resource.kv_extractor import ( - KVRegionExtractorV1, - build_page_table_from_manager, -) -from tensorrt_llm._torch.disaggregation.resource.utils import get_physical_pool, get_pool_bytes +from tensorrt_llm._torch.disaggregation.resource.kv_extractor import KVRegionExtractorV1 +from tensorrt_llm._torch.disaggregation.resource.utils import get_unique_pool_memory_descs from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager -from tensorrt_llm._utils import get_size_in_bytes, nvtx_range +from tensorrt_llm._utils import nvtx_range from tensorrt_llm.disaggregated_params import DisaggregatedParams from tensorrt_llm.runtime.generation import CUASSERT @@ -178,13 +175,11 @@ class Sender(SenderBase): def __init__( self, peer_registrar: PeerRegistrar, - device_id: int, agent: BaseTransferAgent, ): self._registrar = peer_registrar - self._device_id = device_id + self._device_id = peer_registrar.self_rank_info.device_id self._agent = agent - # unique_rid -> instance_rank -> RecvReqInfo self._peer_requests: dict = {} self._peer_requests_lock = threading.Lock() self._messenger = ZMQMessenger(mode="ROUTER") @@ -315,13 +310,15 @@ def _make_agent_request(write_meta: WriteMeta, device_id: int) -> "TransferReque src_dev, dst_dev, mem_type = device_id, write_meta.dst_device_id, MemoryType.VRAM src_list = [ - (ptr, size, src_dev) for ptr, size in zip(write_meta.src_ptrs, write_meta.sizes) + MemoryDesc(ptr, size, src_dev) + for ptr, size in zip(write_meta.src_ptrs, write_meta.sizes) ] dst_list = [ - (ptr, size, dst_dev) for ptr, size in zip(write_meta.dst_ptrs, write_meta.sizes) + MemoryDesc(ptr, size, dst_dev) + for ptr, size in zip(write_meta.dst_ptrs, write_meta.sizes) ] return TransferRequest( - TransferOp.WRITE, + TransferOp.WRITE, # type: ignore[arg-type] MemoryDescs(mem_type, src_list), MemoryDescs(mem_type, dst_list), write_meta.peer_name, @@ -482,12 +479,10 @@ def _build_kv_write_meta(self, task: KVSendTask, req_info: RecvReqInfo) -> Write peer_extractor = self._registrar.peer_extractor( peer_ri.instance_name, peer_ri.instance_rank ) - # Get pool mapping: (self_lg, self_pi) -> (peer_lg, peer_pi) pool_mapping = self._registrar.get_pool_mapping(peer_ri) 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 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] @@ -684,7 +679,6 @@ def has_all_peer_req_infos(self, unique_rid: int) -> bool: return False def _has_all_peer_req_infos(self, req_info: RecvReqInfo) -> bool: - """Checks if all peer info for the request are ready.""" peer_ri = self._registrar.get_peer_rank_info(req_info.instance_name, req_info.instance_rank) expected_transfers = len(self._registrar.get_peer_overlap(peer_ri, peer_ri.dp_rank).ranks) return self._is_req_ready(req_info.unique_rid, expected_transfers) @@ -700,7 +694,6 @@ def shutdown(self): return self._shutdown = True - # Stop all worker threads by sending None to each queue for q in self._send_task_queues: q.put(None) for t in self._worker_threads: @@ -741,14 +734,13 @@ def __init__( request_id: int, params: DisaggregatedParams, sender: Sender, - aux_slot: Optional[int], aux_buffer: Optional[AuxBuffer] = None, ): super().__init__(sender, SessionArgsBase(params)) self._sender: Sender # narrow base class type for Pylance self.request_id = request_id - self.aux_slot = aux_slot self._aux_buffer = aux_buffer + self.aux_slot = aux_buffer.alloc_slot().id if aux_buffer is not None else None self.receiver_ready: bool = False self.kv_tasks = [] self.aux_task = None @@ -849,6 +841,9 @@ def close(self): if getattr(self, "_closed", False): return self._closed = True + if self._aux_buffer is not None and self.aux_slot is not None: + self._aux_buffer.free_slot(self.aux_slot) + self.aux_slot = None # Unregister from Sender; do not null out fields — worker threads # may still access kv_tasks/aux_task/_sender for in-flight transfers. if self._sender is not None: @@ -906,11 +901,9 @@ class Receiver(ReceiverBase): def __init__( self, peer_registrar: PeerRegistrar, - device_id: int, agent: BaseTransferAgent, ): self._registrar = peer_registrar - self._device_id = device_id self._agent = agent self._dealers = {} self._sender_ep_instance_map = {} @@ -1123,14 +1116,13 @@ def __init__( request_id: int, params: DisaggregatedParams, receiver: Receiver, - aux_slot: Optional[int], aux_buffer: Optional[AuxBuffer] = None, ): super().__init__(receiver, SessionArgsBase(params)) self._receiver: Receiver # narrow base class type for Pylance self.request_id = request_id - self.aux_slot = aux_slot self._aux_buffer = aux_buffer + self.aux_slot = aux_buffer.alloc_slot().id if aux_buffer is not None else None self._exception: Optional[Exception] = None self._closed = False self._kv_tasks: list[KVRecvTask] = [] @@ -1277,6 +1269,9 @@ def close(self): if getattr(self, "_closed", False): return self._closed = True + if self._aux_buffer is not None and self.aux_slot is not None: + self._aux_buffer.free_slot(self.aux_slot) + self.aux_slot = None # Unregister from Receiver; do not null out fields — listener thread # may still access _kv_tasks/_receiver for in-flight status messages. if self._receiver is not None: @@ -1351,6 +1346,29 @@ def __exit__(self, _exc_type, _exc_val, _exc_tb): self.shutdown() +def _create_nixl_agent(name: str) -> NixlTransferAgent: + num_threads = int(os.environ.get("TRTLLM_NIXL_NUM_THREADS", "8")) + kwargs = {} + if "TRTLLM_NIXL_SPLIT_BATCH_SIZE" in os.environ: + kwargs["split_batch_size"] = int(os.environ["TRTLLM_NIXL_SPLIT_BATCH_SIZE"]) + return NixlTransferAgent(name, True, num_threads=num_threads, **kwargs) + + +def _make_aux_buffer( + kvm: KVCacheManager, max_slots: int, max_draft_len: Optional[int] = None +) -> Optional[AuxBuffer]: + if max_slots <= 0: + return None + if max_draft_len is None: + max_draft_len = max(0, int(getattr(kvm, "max_draft_len", 0))) + return AuxBuffer( + max_slot_num=max_slots, + beam_width=max(1, int(getattr(kvm, "max_beam_width", 1))), + max_draft_len=max_draft_len, + device="cpu", + ) + + def _deregister_registered_memory(transfer_agent, registered_memorys): try: if transfer_agent is None or not registered_memorys: @@ -1367,60 +1385,29 @@ def _deregister_registered_memory(transfer_agent, registered_memorys): logger.error("unexpected error in _deregister_registered_memory finalizer") -class TransferWorker: - def __init__( - self, - kv_cache_manager: KVCacheManager, - mapping: Mapping, - device_id: int, - instance_name: str, - aux_buffer: Optional[AuxBuffer] = None, - ): - self._mapping = mapping - - self._rank_info: Optional[RankInfo] = None - self._kv_cache_manager = kv_cache_manager - self._aux_buffer = aux_buffer - self._device_id = device_id - self._finalizer = None +@dataclass +class TransferWorkerConfig: + kv_cache_manager: KVCacheManager + device_id: int + instance_name: str + max_concurrent_sessions: int = 0 + max_draft_len: Optional[int] = None - self.init_rank_info(instance_name) - assert self._rank_info is not None - is_leader = self._mapping.rank == 0 - if is_leader: - self._rank_info_server = RankInfoServer(self._rank_info) - else: - self._rank_info_server = None - self._kv_extractor = KVRegionExtractorV1(self._kv_cache_manager) - self._peer_registrar = PeerRegistrar(self._rank_info, self._kv_extractor) - # NixlTransferAgent env config: num_threads for large batches, - # split_batch_size threshold to use dedicated threads (default 1024) - nixl_num_threads = int(os.environ.get("TRTLLM_NIXL_NUM_THREADS", "8")) - nixl_agent_kwargs = {} - if "TRTLLM_NIXL_SPLIT_BATCH_SIZE" in os.environ: - nixl_agent_kwargs["split_batch_size"] = int(os.environ["TRTLLM_NIXL_SPLIT_BATCH_SIZE"]) - - self._agent = NixlTransferAgent( - self._rank_info.instance_name + str(self._rank_info.instance_rank), - True, - num_threads=nixl_num_threads, - **nixl_agent_kwargs, +class TransferWorker: + def __init__(self, config: TransferWorkerConfig): + kvm = config.kv_cache_manager + self._aux_buffer = _make_aux_buffer( + kvm, config.max_concurrent_sessions, config.max_draft_len ) - self._registered_mem = [] - self._register_kv_cache() - if self._aux_buffer is not None: - self._register_aux_buffer() - - self._sender = Sender(self._peer_registrar, device_id, self._agent) - self._receiver = Receiver(self._peer_registrar, device_id, self._agent) - self._rank_info.transfer_engine_info = bytes(self._agent.get_local_agent_desc()) - self._rank_info.self_endpoint = self._receiver.endpoint - - reg_snapshot = list(self._registered_mem) if self._registered_mem is not None else [] - self._finalizer = weakref.finalize( - self, _deregister_registered_memory, self._agent, reg_snapshot + self._rank_info = RankInfo.from_kv_cache_manager( + config.instance_name, + kvm, + config.device_id, + self._aux_buffer.meta if self._aux_buffer is not None else None, ) + self._setup_peer_infrastructure(kvm) + self._setup_transfer_engine() def populate_instance_and_rank_info(self, endpoints: list[str], layer_num_per_pp: list[int]): assert self._rank_info is not None @@ -1428,113 +1415,54 @@ def populate_instance_and_rank_info(self, endpoints: list[str], layer_num_per_pp self._rank_info.layer_num_per_pp = layer_num_per_pp def create_tx_session(self, request: LlmRequest) -> TxSession: - """ - Create a txSession for the request. - """ - if self._aux_buffer is not None: - aux_slot = self._aux_buffer.alloc_slot().id - else: - aux_slot = None params = request.py_disaggregated_params assert params is not None return TxSession( request_id=request.py_request_id, params=params, sender=self._sender, - aux_slot=aux_slot, aux_buffer=self._aux_buffer, ) def create_rx_session(self, request: LlmRequest) -> RxSession: - """ - Create a rxSession for the request. - """ - if self._aux_buffer is not None: - aux_slot = self._aux_buffer.alloc_slot().id - else: - aux_slot = None params = request.py_disaggregated_params assert params is not None return RxSession( request_id=request.py_request_id, params=params, receiver=self._receiver, - aux_slot=aux_slot, aux_buffer=self._aux_buffer, ) - def clear_session(self, session: TxSession | RxSession): - aux_slot = session.aux_slot - if self._aux_buffer is not None: - assert aux_slot is not None - self._aux_buffer.free_slot(aux_slot) - def has_all_peer_req_infos_for_send(self, unique_rid: int) -> bool: return self._sender.has_all_peer_req_infos(unique_rid) - def init_rank_info(self, instance_name): - m = self._mapping - kvm = self._kv_cache_manager - enable_attention_dp = m.enable_attention_dp - - self._rank_info = RankInfo( - instance_name=instance_name, - instance_rank=m.rank, - tp_size=m.tp_size, - tp_rank=m.tp_rank, - pp_size=m.pp_size, - pp_rank=m.pp_rank, - dp_size=m.tp_size if enable_attention_dp else m.dp_size, - dp_rank=m.tp_rank if enable_attention_dp else 0, - cp_size=m.cp_size, - cp_rank=m.cp_rank, - device_id=self._device_id, - layer_num_per_pp=[len(kvm.pp_layers)], - sender_endpoints=[], - server_endpoint="", - self_endpoint="", - transfer_engine_info=bytes(), - attention=AttentionInfo( - kv_heads_per_rank=kvm.num_kv_heads_per_layer[0], - tokens_per_block=kvm.tokens_per_block, - dims_per_head=kvm.head_dim, - element_bytes=get_size_in_bytes(1, kvm.dtype), - enable_attention_dp=enable_attention_dp, - is_mla=kvm.kv_factor == 1, - ), - aux_meta=self._aux_buffer.meta if self._aux_buffer is not None else None, - # Build page table from manager (supports V1 and V2) - page_table=build_page_table_from_manager(kvm), + def _setup_peer_infrastructure(self, kvm: KVCacheManager): + self._rank_info_server = RankInfoServer(self._rank_info) if kvm.mapping.rank == 0 else None + self._kv_extractor = KVRegionExtractorV1(kvm) + self._peer_registrar = PeerRegistrar(self._rank_info, self._kv_extractor) + + def _setup_transfer_engine(self): + self._agent = _create_nixl_agent( + self._rank_info.instance_name + str(self._rank_info.instance_rank) + ) + self._registered_mem: list = [] + self._register_kv_cache() + if self._aux_buffer is not None: + self._register_aux_buffer() + self._sender = Sender(self._peer_registrar, self._agent) + self._receiver = Receiver(self._peer_registrar, self._agent) + self._rank_info.transfer_engine_info = bytes(self._agent.get_local_agent_desc()) + self._rank_info.self_endpoint = self._receiver.endpoint + self._finalizer = weakref.finalize( + self, _deregister_registered_memory, self._agent, list(self._registered_mem) ) def _register_kv_cache(self): - # Get pool information from page_table (works for V1 and V2) - assert self._rank_info is not None - page_table = self._rank_info.page_table - assert page_table is not None - memory_descs = [] - - # Deduplicate pools (different layer_groups may share the same pool) - unique_pools: dict[tuple[int, int], int] = {} # (ptr, size) -> counter - pool_counter = 0 - - for lg_idx, lg in enumerate(page_table.layer_groups): - for pv in lg.pool_views: - pool = get_physical_pool(page_table, lg_idx, pv.pool_idx) - pool_key = (pool.base_address, get_pool_bytes(pool)) - if pool_key not in unique_pools: - unique_pools[pool_key] = pool_counter - pool_counter += 1 - - for (pool_ptr, pool_size), idx in unique_pools.items(): - memory_desc = ( - pool_ptr, - pool_size, - self._device_id, - f"kv_cache_memory_pool{idx}", - ) - memory_descs.append(memory_desc) - + assert self._rank_info.page_table is not None + memory_descs = get_unique_pool_memory_descs( + self._rank_info.page_table, self._rank_info.device_id + ) if memory_descs: reg_memory_desc = RegMemoryDescs("VRAM", memory_descs) self._agent.register_memory(reg_memory_desc) diff --git a/tensorrt_llm/_torch/disaggregation/resource/utils.py b/tensorrt_llm/_torch/disaggregation/resource/utils.py index ceac27b306a5..e8863e402fd3 100644 --- a/tensorrt_llm/_torch/disaggregation/resource/utils.py +++ b/tensorrt_llm/_torch/disaggregation/resource/utils.py @@ -162,6 +162,30 @@ def get_device_pointer( raise ValueError(f"Buffer not found: local_layer_id={local_layer_id}, role={role}") +# ------------------------------------------------------------------------- +# NIXL memory registration helpers +# ------------------------------------------------------------------------- + + +def get_unique_pool_memory_descs( + page_table: KVCachePageTable, device_id: int +) -> list[tuple[int, int, int, str]]: + """Return deduplicated (ptr, size, device_id, name) tuples for all physical pools.""" + unique_pools: dict[tuple[int, int], int] = {} # (ptr, size) -> index + pool_counter = 0 + for lg_idx, lg in enumerate(page_table.layer_groups): + for pv in lg.pool_views: + pool = get_physical_pool(page_table, lg_idx, pv.pool_idx) + pool_key = (pool.base_address, get_pool_bytes(pool)) + if pool_key not in unique_pools: + unique_pools[pool_key] = pool_counter + pool_counter += 1 + return [ + (pool_ptr, pool_size, device_id, f"kv_cache_memory_pool{idx}") + for (pool_ptr, pool_size), idx in unique_pools.items() + ] + + # ------------------------------------------------------------------------- # KVCachePageTable aggregate helpers # ------------------------------------------------------------------------- diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py new file mode 100644 index 000000000000..eaa4192ced70 --- /dev/null +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -0,0 +1,386 @@ +import uuid +from collections import defaultdict +from itertools import chain +from typing import Any, Callable, Dict, List, cast + +import torch + +import tensorrt_llm +from tensorrt_llm import logger +from tensorrt_llm._torch.disaggregation.base.transfer import KVSlice, WaitResult, get_unique_rid +from tensorrt_llm._torch.disaggregation.native.transfer import TransferWorker, TransferWorkerConfig +from tensorrt_llm._torch.disaggregation.resource.utils import get_global_layer_ids +from tensorrt_llm._torch.distributed.communicator import Distributed +from tensorrt_llm._torch.pyexecutor.kv_cache_transceiver import KvCacheTransceiver +from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm.bindings import LlmRequestState +from tensorrt_llm.bindings.executor import ContextPhaseParams +from tensorrt_llm.disaggregated_params import DisaggScheduleStyle +from tensorrt_llm.llmapi.llm_args import CacheTransceiverConfig +from tensorrt_llm.mapping import Mapping + +CacheTransceiverCpp = tensorrt_llm.bindings.internal.batch_manager.CacheTransceiver +AttentionTypeCpp = tensorrt_llm.bindings.internal.batch_manager.AttentionType +CacheTransBufferManagerCpp = tensorrt_llm.bindings.internal.batch_manager.CacheTransBufferManager +BackendTypeCpp = tensorrt_llm.bindings.executor.CacheTransceiverBackendType + + +def _find_consensus_request_ids(request_ids_all_ranks, sync_size): + frequency_map = defaultdict(int) + consensus = [] + for rid in chain.from_iterable(request_ids_all_ranks): + frequency_map[rid] += 1 + for rid, freq in sorted(frequency_map.items(), key=lambda x: x[1], reverse=True): + if freq == sync_size: + consensus.append(rid) + else: + break + return consensus + + +class KvCacheTransceiverV2(KvCacheTransceiver): + def __init__( + self, + mapping: Mapping, + dist: Distributed, + kv_cache_manager: KVCacheManager, + attention_type: AttentionTypeCpp, + cache_transceiver_config: CacheTransceiverConfig, + ): + self._dist: Distributed = dist + self._kv_cache_manager = kv_cache_manager + self._mapping = mapping + self.kv_transfer_timeout_ms = cache_transceiver_config.kv_transfer_timeout_ms + self._sender_future_timeout_ms = ( + cache_transceiver_config.kv_transfer_sender_future_timeout_ms + ) + self._check_compatible() + + self._device_id = torch.cuda.current_device() + logger.info(f"device_id: {self._device_id} in KvCacheTransceiverV2") + self._instance_name = self._broadcast_instance_name() + self._transfer_worker = TransferWorker( + TransferWorkerConfig( + kv_cache_manager=kv_cache_manager, + device_id=self._device_id, + instance_name=self._instance_name, + # * 2: allow back-to-back batches, one in transferring, one preparing next batch + max_concurrent_sessions=max(1, int(kv_cache_manager.max_batch_size)) * 2, + ) + ) + self._dp_rank = mapping.tp_rank if mapping.enable_attention_dp else 0 + self._context_info_endpoint = self._broadcast_context_endpoint() + self._init_sync_policy() + self._exchange_rank_info() + + self._send_sessions = {} + self._recv_sessions = {} + self._send_reqs = {} + self._recv_reqs = {} + self._wait_reqs = {} + self._page_table = self._transfer_worker.page_table + self._is_v2_manager = hasattr(kv_cache_manager, "kv_cache_map") + + def _broadcast_instance_name(self) -> str: + if self._dist.rank == 0: + name = str(uuid.uuid4()) + self._dist.broadcast(name, 0) + return name + return cast(str, self._dist.broadcast(None, 0)) + + def _broadcast_context_endpoint(self) -> str: + if self._dist.rank == 0: + endpoint = self._transfer_worker.rank_info_server_endpoint or "" + self._dist.broadcast(endpoint, 0) + return endpoint + return cast(str, self._dist.broadcast(None, 0)) + + def _init_sync_policy(self): + m = self._mapping + self._ctx_need_tp_sync = m.tp_size > 1 and not m.enable_attention_dp + self._gen_need_sync = not (m.world_size == 1 or (m.enable_attention_dp and m.pp_size == 1)) + pp_allgather: Callable = getattr(self._dist, "pp_allgather") + self._gen_allgather: Callable = ( + pp_allgather if m.enable_attention_dp else self._dist.allgather + ) + + def _exchange_rank_info(self): + endpoints = cast(list, self._dist.allgather(self._transfer_worker.sender_endpoint)) + layer_num_per_pp = cast( + list, getattr(self._dist, "pp_allgather")(len(self._kv_cache_manager.pp_layers)) + ) + self._transfer_worker.populate_instance_and_rank_info( + endpoints=endpoints, layer_num_per_pp=layer_num_per_pp + ) + logger.info(f"transfer worker ctx_server_endpoints: {endpoints}") + logger.info(f"layer_num_per_pp: {layer_num_per_pp}") + logger.info(f"self._context_info_endpoint: {self._context_info_endpoint}") + + def shutdown(self): + if self._transfer_worker is not None: + self._transfer_worker.shutdown() + + def _get_block_ids(self, req: LlmRequest, group_idx: int, lg) -> list: + if self._is_v2_manager: + kv_cache_map = getattr(self._kv_cache_manager, "kv_cache_map") + return list( + kv_cache_map[req.py_request_id].get_aggregated_page_indices( + group_idx, valid_only=True + ) + ) + else: + first_layer = get_global_layer_ids(lg)[0] + return self._kv_cache_manager.get_batch_cache_indices( + [req.py_request_id], layer_idx=first_layer + )[0] + + def _create_kv_slice(self, req: LlmRequest) -> KVSlice: + tpb = self._kv_cache_manager.tokens_per_block + groups = [] + assert self._page_table is not None + for idx, lg in enumerate(self._page_table.layer_groups): + block_ids = self._get_block_ids(req, idx, lg) + + # Filter to only window-relevant blocks for sliding window layer groups. + # Computes the expected number of non-stale blocks (using the same + # eviction formula as update_resources) and keeps only the tail. + # This works correctly regardless of whether update_resources has + # been called: + # - Pre-eviction: all blocks present → trim to last N. + # - Post-eviction (V2 valid_only=True): stale blocks already + # removed → len == expected_valid, so the condition is false. + window_size = lg.sliding_window_size + if window_size is not None: + total_blocks = (req.prompt_len + tpb - 1) // tpb + stale_end = max(0, (req.prompt_len + 1 - window_size) // tpb) + expected_valid = total_blocks - stale_end + if expected_valid <= 0: + block_ids = [] + elif len(block_ids) > expected_valid: + block_ids = block_ids[-expected_valid:] + + groups.append(list(block_ids)) + + return KVSlice(is_last_slice=True, block_ids_per_layer_groups=groups) + + @staticmethod + def _need_aux_transfer(req: LlmRequest) -> bool: + params = req.py_disaggregated_params + return params is not None and params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST + + def _ctx_consensus(self, local_ids: list) -> list: + sync_size = self._dist.tp_size if self._ctx_need_tp_sync else 1 + all_ranks = self._dist.tp_allgather(local_ids) if self._ctx_need_tp_sync else [local_ids] + return _find_consensus_request_ids(all_ranks, sync_size) + + def _gen_consensus(self, local_ids: list) -> list: + sync_size = ( + self._mapping.pp_size if self._mapping.enable_attention_dp else self._mapping.world_size + ) + all_ranks = self._gen_allgather(local_ids) if self._gen_need_sync else [local_ids] + return _find_consensus_request_ids(all_ranks, sync_size) + + def _collect_done(self, sessions: dict, reqs: dict): + """Scan sessions and return (completed_rids, failed_rids).""" + completed, failed = [], [] + for rid, session in sessions.items(): + if session.is_completed(self._need_aux_transfer(reqs[rid])): + completed.append(rid) + elif session.has_failed(): + failed.append(rid) + return completed, failed + + def _build_to_process( + self, sessions: dict, consensus: list, wait_num: int, block_all: bool + ) -> list: + if block_all: + return list(sessions.keys()) + to_process = consensus + for rid in sessions: + if len(to_process) >= wait_num: + break + if rid not in to_process: + to_process.append(rid) + return to_process + + def _close_failed_sessions(self, sessions: dict, reqs: dict, failed: list): + for rid in failed: + reqs[rid].state = LlmRequestState.DISAGG_TRANS_ERROR + sessions[rid].close() + del reqs[rid] + del sessions[rid] + + def _apply_aux(self, session, req: LlmRequest): + """Unpack aux tokens from session into request's context_phase_params.""" + session.unpack_aux(req) + first_gen_tokens = req.py_first_gen_tokens # type: ignore[attr-defined] + draft_tokens = req.py_draft_tokens + if req.context_phase_params is None: + assert req.py_request_id is not None + req.context_phase_params = ContextPhaseParams( + first_gen_tokens=first_gen_tokens, + req_id=req.py_request_id, + opaque_state=b"", + draft_tokens=draft_tokens, + ctx_dp_rank=0, + disagg_info_endpoint="", + ) + else: + req.context_phase_params.first_gen_tokens = first_gen_tokens + req.context_phase_params.draft_tokens = draft_tokens + + def respond_and_send_async(self, req: LlmRequest): + rid = get_unique_rid(req) + assert rid is not None + if rid not in self._send_sessions: + self._send_sessions[rid] = self._transfer_worker.create_tx_session(req) + session = self._send_sessions[rid] + req.state = LlmRequestState.DISAGG_CONTEXT_TRANS_IN_PROGRESS + session.send(self._create_kv_slice(req)) + if self._need_aux_transfer(req): + session.pack_aux(req) + session.send_aux() + req.context_phase_params = ContextPhaseParams( + first_gen_tokens=[], + req_id=rid, + opaque_state=None, + draft_tokens=None, + ctx_dp_rank=self._dp_rank, + disagg_info_endpoint=self._context_info_endpoint, + ) + self._send_reqs[rid] = req + + def request_and_receive_sync(self, req: LlmRequest): + raise NotImplementedError("request_and_receive_sync is not implemented") + + def request_and_receive_async(self, req: LlmRequest): + rid = get_unique_rid(req) + 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)) + self._recv_reqs[rid] = req + + def check_context_transfer_status(self, at_least_request_num: int, mark_complete: bool = False): + block_all = at_least_request_num is None + wait_num = at_least_request_num if not block_all else 0 + + local_completed, local_failed = self._collect_done(self._send_sessions, self._send_reqs) + to_process = self._build_to_process( + self._send_sessions, + self._ctx_consensus(local_completed + local_failed), + wait_num, + block_all, + ) + + completed, timed_out, failed = [], [], [] + timeout = self._sender_future_timeout_ms / 1000.0 + for rid in to_process: + session = self._send_sessions[rid] + result = session.wait_complete( + need_aux=self._need_aux_transfer(self._send_reqs[rid]), timeout=timeout + ) + if result == WaitResult.COMPLETED: + completed.append(rid) + elif result == WaitResult.TIMEOUT: + logger.warning( + f"TxSession rid={session.disagg_request_id} timed out after {self._sender_future_timeout_ms}ms" + ) + timed_out.append(rid) + else: + logger.warning(f"TxSession rid={session.disagg_request_id} failed") + failed.append(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) + + return completed, failed + + def check_gen_transfer_status(self, at_least_request_num: int): + block_all = at_least_request_num is None + wait_num = at_least_request_num if not block_all else 0 + + local_completed, local_failed = self._collect_done(self._recv_sessions, self._recv_reqs) + to_process = self._build_to_process( + self._recv_sessions, + self._gen_consensus(local_completed + local_failed), + wait_num, + block_all, + ) + + completed, failed = [], [] + for rid in to_process: + result = self._recv_sessions[rid].wait_complete( + need_aux=self._need_aux_transfer(self._recv_reqs[rid]), block_for_aux=block_all + ) + if result == WaitResult.COMPLETED: + completed.append(rid) + elif result == WaitResult.FAILED: + failed.append(rid) + # else: None — KV done but aux still in flight; re-poll next cycle + + for rid in completed: + session = self._recv_sessions[rid] + req = self._recv_reqs[rid] + if self._need_aux_transfer(req): + self._apply_aux(session, req) + req.state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE + session.close() + del self._recv_reqs[rid] + del self._recv_sessions[rid] + self._close_failed_sessions(self._recv_sessions, self._recv_reqs, failed) + + def check_gen_transfer_complete(self): + return len(self._recv_sessions) == 0 + + def cancel_request(self, req: LlmRequest): + raise NotImplementedError("cancel_request is not implemented") + + def get_disaggregated_params(self) -> Dict[str, Any]: + # Keep this aligned with fields populated in respond_and_send_async(). + # These values are server-level metadata used to seed generation-first + # requests before context-phase response data arrives. + return { + "ctx_dp_rank": self._dp_rank, + "ctx_info_endpoint": [self._context_info_endpoint] + if self._context_info_endpoint + else None, + } + + def prepare_context_requests(self, requests: List[LlmRequest]): + # Place new generation-first context requests into wait state, then + # use tp_allgather consensus to promote ready requests to CONTEXT_INIT. + for req in requests: + rid = get_unique_rid(req) + if rid not in self._send_sessions: + self._wait_reqs[rid] = req + req.state = LlmRequestState.DISAGG_CONTEXT_WAIT_SCHEDULER + + # Check which waiting requests have peer info locally, then tp_allgather + # consensus so all TP ranks agree before promoting. + # Without consensus, background peer info arriving at different times on + # different ranks causes scheduling mismatches → hang. + local_ready = [ + rid + for rid in self._wait_reqs + if self._transfer_worker.has_all_peer_req_infos_for_send(rid) + ] + for rid in self._ctx_consensus(local_ready): + self._wait_reqs[rid].state = LlmRequestState.CONTEXT_INIT + del self._wait_reqs[rid] + + def _check_compatible(self): + if self._mapping.cp_size != 1: + raise ValueError( + f"KvCacheTransceiverV2: _check_compatible: only support context parallelism is 1: " + f"cp_size: {self._mapping.cp_size}" + ) + + def get_context_state(self): + raise NotImplementedError("get_context_state is not implemented") diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py index 8716590d86cb..df4a11be9016 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py @@ -80,12 +80,11 @@ def create_kv_cache_transceiver( f"got {cache_transceiver_config.backend}. " f"Please use transceiver_runtime='CPP' for MPI, UCX, or MOONCAKE backends." ) - from tensorrt_llm._torch.disaggregation.native.py_cache_transceiver import \ - PyNativeCacheTransceiver - logger.info("Using PyNativeCacheTransceiver") - return PyNativeCacheTransceiver(mapping, dist, kv_cache_manager, - attention_type, - cache_transceiver_config) + from tensorrt_llm._torch.disaggregation.transceiver import \ + KvCacheTransceiverV2 + logger.info("Using KvCacheTransceiverV2") + return KvCacheTransceiverV2(mapping, dist, kv_cache_manager, + attention_type, cache_transceiver_config) # Default: use C++ transceiver (transceiver_runtime is None or "CPP") return BindKvCacheTransceiver(mapping, dist, kv_cache_manager, diff --git a/tests/unittest/disaggregated/test_kv_transfer.py b/tests/unittest/disaggregated/test_kv_transfer.py index 6347db2d6c09..73f7a4411841 100644 --- a/tests/unittest/disaggregated/test_kv_transfer.py +++ b/tests/unittest/disaggregated/test_kv_transfer.py @@ -20,8 +20,7 @@ SessionStatus, TokenRange, ) -from tensorrt_llm._torch.disaggregation.native.auxiliary import AuxBuffer -from tensorrt_llm._torch.disaggregation.native.transfer import TransferWorker +from tensorrt_llm._torch.disaggregation.native.transfer import TransferWorker, TransferWorkerConfig from tensorrt_llm._torch.disaggregation.resource.kv_extractor import KVRegionExtractorV1 from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, LlmRequestType from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager, KVCacheManagerV2 @@ -175,10 +174,6 @@ def create_transfer_worker_setup( ) ) - meta_max_batch_size = 32 - beam_width = 1 - max_draft_len = 4 - ctx_instance_num = ctx_tp * ctx_pp gen_instance_num = gen_tp * gen_pp num_layers = 4 @@ -202,7 +197,6 @@ def create_transfer_worker_setup( request_len = 16 for i in range(ctx_instance_num): - ctx_aux_buffer = AuxBuffer(meta_max_batch_size, beam_width, max_draft_len) cache_type = ( tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF if not is_mla @@ -320,11 +314,13 @@ def create_transfer_worker_setup( ctx_kv_cache_managers.append(ctx_kv_cache_manager) ctx_transfer_workers.append( TransferWorker( - kv_cache_manager=ctx_kv_cache_manager, - mapping=ctx_mappings[i], - device_id=device_id, - instance_name=ctx_instance_name, - aux_buffer=ctx_aux_buffer, + TransferWorkerConfig( + kv_cache_manager=ctx_kv_cache_manager, + device_id=device_id, + instance_name=ctx_instance_name, + max_concurrent_sessions=max_batch_size * 2, + max_draft_len=4, + ) ) ) @@ -334,9 +330,7 @@ def create_transfer_worker_setup( ] ctx_layer_num_per_pp = [] for pp_rank in range(ctx_pp): - ctx_layer_num_per_pp.append( - len(ctx_transfer_workers[pp_rank * ctx_tp]._kv_cache_manager.pp_layers) - ) + ctx_layer_num_per_pp.append(len(ctx_kv_cache_managers[pp_rank * ctx_tp].pp_layers)) for ctx_transfer_worker in ctx_transfer_workers: ctx_transfer_worker.populate_instance_and_rank_info( @@ -346,7 +340,6 @@ def create_transfer_worker_setup( gen_transfer_workers = [] gen_kv_cache_managers = [] for i in range(gen_instance_num): - gen_aux_buffer = AuxBuffer(meta_max_batch_size, beam_width, max_draft_len) cache_type = ( tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF if not is_mla @@ -437,11 +430,13 @@ def create_transfer_worker_setup( gen_kv_cache_managers.append(gen_kv_cache_manager) gen_transfer_workers.append( TransferWorker( - kv_cache_manager=gen_kv_cache_manager, - mapping=gen_mappings[i], - device_id=device_id, - instance_name=gen_instance_name, - aux_buffer=gen_aux_buffer, + TransferWorkerConfig( + kv_cache_manager=gen_kv_cache_manager, + device_id=device_id, + instance_name=gen_instance_name, + max_concurrent_sessions=max_batch_size * 2, + max_draft_len=4, + ) ) ) _ = gen_transfer_workers[0]._rank_info_server.endpoint # noqa: F841 @@ -450,9 +445,7 @@ def create_transfer_worker_setup( ] gen_layer_num_per_pp = [] for pp_rank in range(gen_pp): - gen_layer_num_per_pp.append( - len(gen_transfer_workers[pp_rank * gen_tp]._kv_cache_manager.pp_layers) - ) + gen_layer_num_per_pp.append(len(gen_kv_cache_managers[pp_rank * gen_tp].pp_layers)) for gen_transfer_worker in gen_transfer_workers: gen_transfer_worker.populate_instance_and_rank_info( endpoints=gen_endpoints, layer_num_per_pp=gen_layer_num_per_pp @@ -899,7 +892,6 @@ def get_layers_in_group_per_pp(kv_cache_managers, pp_size, tp_size, group_id, is if not send_first: for pp_rank in range(gen_pp): for tp_rank in range(valid_gen_tp): - transfer_worker = valid_gen_transfer_workers[pp_rank * valid_gen_tp + tp_rank] recv_session = receiver_sessions[pp_rank * valid_gen_tp + tp_rank] recv_session.unpack_aux(gen_request) @@ -910,10 +902,10 @@ def get_layers_in_group_per_pp(kv_cache_managers, pp_size, tp_size, group_id, is 11 + ctx_request_id, 12 + ctx_request_id, ] - for transfer_worker, receiver_session in zip(valid_gen_transfer_workers, receiver_sessions): - transfer_worker.clear_session(receiver_session) - for transfer_worker, sender_session in zip(valid_ctx_transfer_workers, sender_sessions): - transfer_worker.clear_session(sender_session) + for receiver_session in receiver_sessions: + receiver_session.close() + for sender_session in sender_sessions: + sender_session.close() # V2: Close kv_caches to release slots if use_v2: diff --git a/tests/unittest/disaggregated/test_kv_transfer_mp.py b/tests/unittest/disaggregated/test_kv_transfer_mp.py index 55455b815d85..d4f12e97da71 100644 --- a/tests/unittest/disaggregated/test_kv_transfer_mp.py +++ b/tests/unittest/disaggregated/test_kv_transfer_mp.py @@ -11,8 +11,7 @@ import tensorrt_llm.bindings.executor as trtllm from tensorrt_llm import DisaggregatedParams, Mapping, SamplingParams from tensorrt_llm._torch.disaggregation.base.transfer import KVSlice, SessionStatus -from tensorrt_llm._torch.disaggregation.native.auxiliary import AuxBuffer -from tensorrt_llm._torch.disaggregation.native.transfer import TransferWorker +from tensorrt_llm._torch.disaggregation.native.transfer import TransferWorker, TransferWorkerConfig from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest, LlmRequestType from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm.bindings import DataType @@ -139,9 +138,6 @@ def worker_fn( ctx_to_gen_group = dist.new_group(ranks=[0] + gen_ranks) # Common parameters - meta_max_batch_size = 32 - beam_width = 1 - max_draft_len = 4 num_layers = 4 head_dim = 128 num_kv_heads = 4 @@ -187,13 +183,13 @@ def worker_fn( block_data_pool.copy_(random_values) # Create TransferWorker - aux_buffer = AuxBuffer(meta_max_batch_size, beam_width, max_draft_len) transfer_worker = TransferWorker( - kv_cache_manager=kv_cache_manager, - mapping=mapping, - device_id=device_id, - instance_name=ctx_instance_name, - aux_buffer=aux_buffer, + TransferWorkerConfig( + kv_cache_manager=kv_cache_manager, + device_id=device_id, + instance_name=ctx_instance_name, + max_concurrent_sessions=max_batch_size * 2, + ) ) # Get local endpoint @@ -249,13 +245,13 @@ def worker_fn( ) # Create TransferWorker - aux_buffer = AuxBuffer(meta_max_batch_size, beam_width, max_draft_len) transfer_worker = TransferWorker( - kv_cache_manager=kv_cache_manager, - mapping=mapping, - device_id=device_id, - instance_name=gen_instance_name, - aux_buffer=aux_buffer, + TransferWorkerConfig( + kv_cache_manager=kv_cache_manager, + device_id=device_id, + instance_name=gen_instance_name, + max_concurrent_sessions=max_batch_size * 2, + ) ) # Get local endpoint diff --git a/tests/unittest/disaggregated/test_py_cache_transceiver_mp.py b/tests/unittest/disaggregated/test_py_cache_transceiver_mp.py index 24099de4b7ee..f77ada82a529 100644 --- a/tests/unittest/disaggregated/test_py_cache_transceiver_mp.py +++ b/tests/unittest/disaggregated/test_py_cache_transceiver_mp.py @@ -1,4 +1,4 @@ -"""Multi-process test for PyNativeCacheTransceiver (V2 backend). +"""Multi-process test for KvCacheTransceiverV2 (V2 backend). This test uses torch.multiprocessing to spawn multiple processes simulating ctx and gen instances with different TP/PP configurations. @@ -105,7 +105,7 @@ def find_free_port(): class TorchDistributedWrapper: """A wrapper that provides the Distributed interface using torch.distributed. - This is used to create a compatible Distributed object for PyNativeCacheTransceiver + This is used to create a compatible Distributed object for KvCacheTransceiverV2 in multi-process tests. """ @@ -289,10 +289,8 @@ def on_hang_detected(): else tensorrt_llm.bindings.internal.batch_manager.CacheType.SELF ) - # Import PyNativeCacheTransceiver - from tensorrt_llm._torch.disaggregation.native.py_cache_transceiver import ( - PyNativeCacheTransceiver, - ) + # Import KvCacheTransceiverV2 + from tensorrt_llm._torch.disaggregation.transceiver import KvCacheTransceiverV2 # ===== Create all TP/PP groups in the same order for all ranks ===== # dist.new_group is a collective operation - ALL ranks must call it in the same order! @@ -398,10 +396,10 @@ def on_hang_detected(): backend="NIXL", transceiver_runtime="PYTHON", max_tokens_in_buffer=512 ) - # Create PyNativeCacheTransceiver + # Create KvCacheTransceiverV2 attention_type = AttentionTypeCpp.MLA if is_mla else AttentionTypeCpp.DEFAULT print(f"[Rank {rank}] CTX: Creating transceiver...", flush=True) - transceiver = PyNativeCacheTransceiver( + transceiver = KvCacheTransceiverV2( mapping=mapping, dist=dist_wrapper, kv_cache_manager=kv_cache_manager, @@ -409,7 +407,8 @@ def on_hang_detected(): cache_transceiver_config=cache_transceiver_config, ) print(f"[Rank {rank}] CTX: Transceiver created", flush=True) - ctx_info_endpoint = transceiver.context_info_endpoint if local_rank == 0 else None + endpoints = transceiver.get_disaggregated_params().get("ctx_info_endpoint") or [] + ctx_info_endpoint = endpoints[0] if (local_rank == 0 and endpoints) else None else: # gen process # Create gen mapping @@ -466,10 +465,10 @@ def on_hang_detected(): backend="NIXL", transceiver_runtime="PYTHON", max_tokens_in_buffer=512 ) - # Create PyNativeCacheTransceiver + # Create KvCacheTransceiverV2 attention_type = AttentionTypeCpp.MLA if is_mla else AttentionTypeCpp.DEFAULT print(f"[Rank {rank}] GEN: Creating transceiver...", flush=True) - transceiver = PyNativeCacheTransceiver( + transceiver = KvCacheTransceiverV2( mapping=mapping, dist=dist_wrapper, kv_cache_manager=kv_cache_manager, @@ -1024,7 +1023,7 @@ def run_v2_transceiver_mp( is_mla: bool = False, ctx_gen_workflow: str = "ctx_first", ): - """Multi-process test for PyNativeCacheTransceiver using mp.spawn.""" + """Multi-process test for KvCacheTransceiverV2 using mp.spawn.""" world_size = ctx_tp * ctx_pp + gen_tp * gen_pp master_addr = "127.0.0.1" @@ -1099,7 +1098,7 @@ def run_v2_transceiver_mp( def test_v2_transceiver_mp( ctx_tp, ctx_pp, gen_tp, gen_pp, ctx_enable_dp, gen_enable_dp, is_mla, workflow ): - """Test PyNativeCacheTransceiver with context-first multi-process configurations.""" + """Test KvCacheTransceiverV2 with context-first multi-process configurations.""" try: mp.set_start_method("spawn", force=True) except RuntimeError: From 8d2f15e9dba40d3f4c6a0329f98c8297075a6847 Mon Sep 17 00:00:00 2001 From: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> Date: Mon, 23 Mar 2026 09:20:57 +0000 Subject: [PATCH 2/5] fix according to comments Signed-off-by: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> --- .../_torch/disaggregation/native/rank_info.py | 3 +- .../_torch/disaggregation/native/transfer.py | 61 +++++++++++-------- .../_torch/disaggregation/transceiver.py | 44 ++++++++----- .../_torch/pyexecutor/kv_cache_transceiver.py | 2 +- 4 files changed, 68 insertions(+), 42 deletions(-) diff --git a/tensorrt_llm/_torch/disaggregation/native/rank_info.py b/tensorrt_llm/_torch/disaggregation/native/rank_info.py index c908b7e415da..12a614ca2dad 100644 --- a/tensorrt_llm/_torch/disaggregation/native/rank_info.py +++ b/tensorrt_llm/_torch/disaggregation/native/rank_info.py @@ -7,6 +7,7 @@ from tensorrt_llm._torch.disaggregation.native.mixers.attention.spec import AttentionInfo from tensorrt_llm._torch.disaggregation.resource.kv_extractor import build_page_table_from_manager from tensorrt_llm._torch.disaggregation.resource.page import KVCachePageTable +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm._utils import get_size_in_bytes @@ -51,7 +52,7 @@ def to_bytes(self) -> bytes: def from_kv_cache_manager( cls, instance_name: str, - kv_cache_manager, + kv_cache_manager: KVCacheManager, device_id: int, aux_buffer_meta: Optional[AuxBufferMeta] = None, ) -> "RankInfo": diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index dfca6dab0e9c..0616b881120f 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -299,6 +299,13 @@ def _process_task_queue(self, thread_idx: int): @staticmethod @nvtx_range("_make_agent_request") def _make_agent_request(write_meta: WriteMeta, device_id: int) -> "TransferRequest": + if not (len(write_meta.src_ptrs) == len(write_meta.dst_ptrs) == len(write_meta.sizes)): + raise ValueError( + f"Pointer/size mismatch for unique_rid={write_meta.unique_rid}: " + f"{len(write_meta.src_ptrs)=}, " + f"{len(write_meta.dst_ptrs)=}, " + f"{len(write_meta.sizes)=}" + ) if write_meta.meta_type == WriteMetaType.AUX: src_dev, dst_dev, mem_type = 0, 0, MemoryType.DRAM else: @@ -311,11 +318,11 @@ def _make_agent_request(write_meta: WriteMeta, device_id: int) -> "TransferReque src_list = [ MemoryDesc(ptr, size, src_dev) - for ptr, size in zip(write_meta.src_ptrs, write_meta.sizes) + for ptr, size in zip(write_meta.src_ptrs, write_meta.sizes, strict=True) ] dst_list = [ MemoryDesc(ptr, size, dst_dev) - for ptr, size in zip(write_meta.dst_ptrs, write_meta.sizes) + for ptr, size in zip(write_meta.dst_ptrs, write_meta.sizes, strict=True) ] return TransferRequest( TransferOp.WRITE, # type: ignore[arg-type] @@ -809,9 +816,10 @@ def wait_complete(self, need_aux: bool, timeout: float) -> WaitResult: Returns WaitResult.COMPLETED, WaitResult.FAILED, or WaitResult.TIMEOUT. """ try: - kv_status = self.kv_tasks[0].future.result(timeout=timeout) - if kv_status != AgentResult.SUCCESS: - return WaitResult.FAILED + for task in self.kv_tasks: + kv_status = task.future.result(timeout=timeout) + if kv_status != AgentResult.SUCCESS: + return WaitResult.FAILED if need_aux and self.aux_task is not None: aux_status = self.aux_task.future.result(timeout=timeout) if aux_status != AgentResult.SUCCESS: @@ -1132,17 +1140,15 @@ def __init__( @property def status(self) -> SessionStatus: - if self._exception is not None or ( - self._kv_tasks and self._kv_tasks[0].status == TaskStatus.ERROR - ): + if self._exception is not None or any(t.status == TaskStatus.ERROR for t in self._kv_tasks): return SessionStatus.ERROR if self._kv_tasks: - task_status = self._kv_tasks[0].status - if task_status == TaskStatus.TRANSFERRED and self._aux_status == TaskStatus.TRANSFERRED: + kv_all_transferred = all(t.status == TaskStatus.TRANSFERRED for t in self._kv_tasks) + if kv_all_transferred and self._aux_status == TaskStatus.TRANSFERRED: return SessionStatus.FULLY_TRANSFERRED - if task_status == TaskStatus.TRANSFERRED: + if kv_all_transferred: return SessionStatus.KV_TRANSFERRED - if task_status == TaskStatus.TRANSFERRING: + if any(t.status == TaskStatus.TRANSFERRING for t in self._kv_tasks): return SessionStatus.TRANSFERRING return SessionStatus.INIT @@ -1196,7 +1202,9 @@ def process_kv_agent_result( ) def process_aux_agent_result(self, _peer_rank: int, status: AgentResult): - task = self._kv_tasks[0] # TODO: index by slice_id when multi-slice is supported + # Aux is session-level (not per-slice); expected_transfers is identical + # across all kv_tasks, so any task provides the right count. + task = self._kv_tasks[0] if status == AgentResult.SUCCESS: self._aux_count += 1 @@ -1248,9 +1256,10 @@ def wait_complete(self, need_aux: bool, block_for_aux: bool = False) -> Optional Returns WaitResult.COMPLETED on full success, WaitResult.FAILED on error. """ try: - kv_status = self._kv_tasks[0].future.result() - if kv_status != AgentResult.SUCCESS: - return WaitResult.FAILED + for task in self._kv_tasks: + kv_status = task.future.result() + if kv_status != AgentResult.SUCCESS: + return WaitResult.FAILED if need_aux: while True: status = self.status @@ -1447,16 +1456,20 @@ def _setup_transfer_engine(self): self._rank_info.instance_name + str(self._rank_info.instance_rank) ) self._registered_mem: list = [] - self._register_kv_cache() - if self._aux_buffer is not None: - self._register_aux_buffer() - self._sender = Sender(self._peer_registrar, self._agent) - self._receiver = Receiver(self._peer_registrar, self._agent) - self._rank_info.transfer_engine_info = bytes(self._agent.get_local_agent_desc()) - self._rank_info.self_endpoint = self._receiver.endpoint self._finalizer = weakref.finalize( - self, _deregister_registered_memory, self._agent, list(self._registered_mem) + self, _deregister_registered_memory, self._agent, self._registered_mem ) + try: + self._register_kv_cache() + if self._aux_buffer is not None: + self._register_aux_buffer() + self._sender = Sender(self._peer_registrar, self._agent) + self._receiver = Receiver(self._peer_registrar, self._agent) + self._rank_info.transfer_engine_info = bytes(self._agent.get_local_agent_desc()) + self._rank_info.self_endpoint = self._receiver.endpoint + except Exception: + self._finalizer() + raise def _register_kv_cache(self): assert self._rank_info.page_table is not None diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index eaa4192ced70..69f562cc92d6 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -1,11 +1,10 @@ import uuid from collections import defaultdict from itertools import chain -from typing import Any, Callable, Dict, List, cast +from typing import Any, Callable, Dict, List, Optional, cast import torch -import tensorrt_llm from tensorrt_llm import logger from tensorrt_llm._torch.disaggregation.base.transfer import KVSlice, WaitResult, get_unique_rid from tensorrt_llm._torch.disaggregation.native.transfer import TransferWorker, TransferWorkerConfig @@ -13,18 +12,13 @@ from tensorrt_llm._torch.distributed.communicator import Distributed from tensorrt_llm._torch.pyexecutor.kv_cache_transceiver import KvCacheTransceiver from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest -from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager +from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager, KVCacheManagerV2 from tensorrt_llm.bindings import LlmRequestState from tensorrt_llm.bindings.executor import ContextPhaseParams from tensorrt_llm.disaggregated_params import DisaggScheduleStyle from tensorrt_llm.llmapi.llm_args import CacheTransceiverConfig from tensorrt_llm.mapping import Mapping -CacheTransceiverCpp = tensorrt_llm.bindings.internal.batch_manager.CacheTransceiver -AttentionTypeCpp = tensorrt_llm.bindings.internal.batch_manager.AttentionType -CacheTransBufferManagerCpp = tensorrt_llm.bindings.internal.batch_manager.CacheTransBufferManager -BackendTypeCpp = tensorrt_llm.bindings.executor.CacheTransceiverBackendType - def _find_consensus_request_ids(request_ids_all_ranks, sync_size): frequency_map = defaultdict(int) @@ -45,7 +39,6 @@ def __init__( mapping: Mapping, dist: Distributed, kv_cache_manager: KVCacheManager, - attention_type: AttentionTypeCpp, cache_transceiver_config: CacheTransceiverConfig, ): self._dist: Distributed = dist @@ -80,7 +73,7 @@ def __init__( self._recv_reqs = {} self._wait_reqs = {} self._page_table = self._transfer_worker.page_table - self._is_v2_manager = hasattr(kv_cache_manager, "kv_cache_map") + self._is_v2_manager = isinstance(kv_cache_manager, KVCacheManagerV2) def _broadcast_instance_name(self) -> str: if self._dist.rank == 0: @@ -113,13 +106,23 @@ def _exchange_rank_info(self): self._transfer_worker.populate_instance_and_rank_info( endpoints=endpoints, layer_num_per_pp=layer_num_per_pp ) - logger.info(f"transfer worker ctx_server_endpoints: {endpoints}") + logger.info(f"transfer worker ctx_server_endpoints: {endpoints}") logger.info(f"layer_num_per_pp: {layer_num_per_pp}") logger.info(f"self._context_info_endpoint: {self._context_info_endpoint}") def shutdown(self): - if self._transfer_worker is not None: - self._transfer_worker.shutdown() + if getattr(self, "_shutdown", False): + return + self._shutdown = True + for session in list(self._send_sessions.values()): + session.close() + for session in list(self._recv_sessions.values()): + session.close() + self._send_sessions.clear() + self._send_reqs.clear() + self._recv_sessions.clear() + self._recv_reqs.clear() + self._transfer_worker.shutdown() def _get_block_ids(self, req: LlmRequest, group_idx: int, lg) -> list: if self._is_v2_manager: @@ -196,7 +199,7 @@ def _build_to_process( ) -> list: if block_all: return list(sessions.keys()) - to_process = consensus + to_process = list(consensus) for rid in sessions: if len(to_process) >= wait_num: break @@ -256,13 +259,20 @@ def request_and_receive_sync(self, req: LlmRequest): def request_and_receive_async(self, req: LlmRequest): rid = get_unique_rid(req) + if rid in self._recv_sessions: + logger.warning( + f"request_and_receive_async: rid={rid} already has a recv session, skipping" + ) + return 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)) self._recv_reqs[rid] = req - def check_context_transfer_status(self, at_least_request_num: int, mark_complete: bool = False): + def check_context_transfer_status( + self, at_least_request_num: Optional[int], mark_complete: bool = False + ): block_all = at_least_request_num is None wait_num = at_least_request_num if not block_all else 0 @@ -302,7 +312,7 @@ def check_context_transfer_status(self, at_least_request_num: int, mark_complete return completed, failed - def check_gen_transfer_status(self, at_least_request_num: int): + def check_gen_transfer_status(self, at_least_request_num: Optional[int]): block_all = at_least_request_num is None wait_num = at_least_request_num if not block_all else 0 @@ -336,6 +346,8 @@ def check_gen_transfer_status(self, at_least_request_num: int): del self._recv_sessions[rid] self._close_failed_sessions(self._recv_sessions, self._recv_reqs, failed) + return completed, failed + def check_gen_transfer_complete(self): return len(self._recv_sessions) == 0 diff --git a/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py b/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py index df4a11be9016..c84fc8f5d637 100644 --- a/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py +++ b/tensorrt_llm/_torch/pyexecutor/kv_cache_transceiver.py @@ -84,7 +84,7 @@ def create_kv_cache_transceiver( KvCacheTransceiverV2 logger.info("Using KvCacheTransceiverV2") return KvCacheTransceiverV2(mapping, dist, kv_cache_manager, - attention_type, cache_transceiver_config) + cache_transceiver_config) # Default: use C++ transceiver (transceiver_runtime is None or "CPP") return BindKvCacheTransceiver(mapping, dist, kv_cache_manager, From 5d2b9261472ed5e23088a369a20b12b25693dbed Mon Sep 17 00:00:00 2001 From: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> Date: Mon, 23 Mar 2026 09:40:40 +0000 Subject: [PATCH 3/5] increase the max sessions num Signed-off-by: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> --- tensorrt_llm/_torch/disaggregation/transceiver.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index 69f562cc92d6..bcf3f2e75c09 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -58,8 +58,10 @@ def __init__( kv_cache_manager=kv_cache_manager, device_id=self._device_id, instance_name=self._instance_name, - # * 2: allow back-to-back batches, one in transferring, one preparing next batch - max_concurrent_sessions=max(1, int(kv_cache_manager.max_batch_size)) * 2, + # Context-only requests are released after KV transfer completes, so many batches + # can be in-flight simultaneously. AuxBuffer holds only small CPU metadata, so a + # large multiplier is cheap. + max_concurrent_sessions=max(1, int(kv_cache_manager.max_batch_size)) * 20000, ) ) self._dp_rank = mapping.tp_rank if mapping.enable_attention_dp else 0 From c3a682b43afebf085b9e3aa27123720512de7a7c Mon Sep 17 00:00:00 2001 From: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> Date: Mon, 23 Mar 2026 12:38:52 +0000 Subject: [PATCH 4/5] fix the test Signed-off-by: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> --- .../unittest/disaggregated/test_py_cache_transceiver_mp.py | 6 ------ 1 file changed, 6 deletions(-) diff --git a/tests/unittest/disaggregated/test_py_cache_transceiver_mp.py b/tests/unittest/disaggregated/test_py_cache_transceiver_mp.py index f77ada82a529..64992bfba688 100644 --- a/tests/unittest/disaggregated/test_py_cache_transceiver_mp.py +++ b/tests/unittest/disaggregated/test_py_cache_transceiver_mp.py @@ -26,8 +26,6 @@ from tensorrt_llm.disaggregated_params import DisaggScheduleStyle from tensorrt_llm.llmapi.llm_args import CacheTransceiverConfig -AttentionTypeCpp = tensorrt_llm.bindings.internal.batch_manager.AttentionType - def broadcast_string(s: str | None, src: int, group: dist.ProcessGroup | None = None) -> str: """Broadcast a string from src rank to all other ranks in the group.""" @@ -397,13 +395,11 @@ def on_hang_detected(): ) # Create KvCacheTransceiverV2 - attention_type = AttentionTypeCpp.MLA if is_mla else AttentionTypeCpp.DEFAULT print(f"[Rank {rank}] CTX: Creating transceiver...", flush=True) transceiver = KvCacheTransceiverV2( mapping=mapping, dist=dist_wrapper, kv_cache_manager=kv_cache_manager, - attention_type=attention_type, cache_transceiver_config=cache_transceiver_config, ) print(f"[Rank {rank}] CTX: Transceiver created", flush=True) @@ -466,13 +462,11 @@ def on_hang_detected(): ) # Create KvCacheTransceiverV2 - attention_type = AttentionTypeCpp.MLA if is_mla else AttentionTypeCpp.DEFAULT print(f"[Rank {rank}] GEN: Creating transceiver...", flush=True) transceiver = KvCacheTransceiverV2( mapping=mapping, dist=dist_wrapper, kv_cache_manager=kv_cache_manager, - attention_type=attention_type, cache_transceiver_config=cache_transceiver_config, ) print(f"[Rank {rank}] GEN: Transceiver created", flush=True) From 9ba574dbd2722ab0e3c5c9521388b923fef71fad Mon Sep 17 00:00:00 2001 From: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> Date: Tue, 24 Mar 2026 10:49:11 +0000 Subject: [PATCH 5/5] refactor the interfaces Signed-off-by: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> --- .../_torch/disaggregation/base/transfer.py | 78 ++++++------------- .../_torch/disaggregation/native/transfer.py | 45 +++++++---- .../_torch/disaggregation/transceiver.py | 25 +++--- 3 files changed, 68 insertions(+), 80 deletions(-) diff --git a/tensorrt_llm/_torch/disaggregation/base/transfer.py b/tensorrt_llm/_torch/disaggregation/base/transfer.py index d7e2fc071093..842530620cfa 100644 --- a/tensorrt_llm/_torch/disaggregation/base/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/base/transfer.py @@ -4,7 +4,7 @@ from abc import ABC, abstractmethod from dataclasses import dataclass, field from enum import Enum -from typing import List, Optional +from typing import List, Optional, cast from tensorrt_llm import DisaggregatedParams from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest @@ -108,74 +108,46 @@ class ReceiverBase(ABC): ... -class TxSessionBase(ABC): - def __init__(self, sender: SenderBase, args: SessionArgsBase): - """ - Initializes the transmission session. - :param sender: The sender instance responsible for sending data. - :param args: The session arguments. - """ - self._sender = sender +class _SessionBase(ABC): + """Shared base for Tx/Rx sessions.""" + + def __init__(self, args: SessionArgsBase): self._base_args = args @property def disagg_request_id(self) -> int: - return self._base_args.params.disagg_request_id + return cast(int, self._base_args.params.disagg_request_id) @abstractmethod - def send(self, slice: KVSlice) -> concurrent.futures.Future: - """ - Sends a slice of KV cache data and returns a Future for the transfer. - :param slice: The KV slice to send. - """ - ... + def is_completed(self) -> bool: ... - @property @abstractmethod - def exception(self) -> Optional[Exception]: - """ - Returns any exception that occurred during the session. - """ - ... + def wait_complete(self) -> Optional[WaitResult]: ... + @property @abstractmethod - def close(self) -> None: - """ - Closes the session and releases any resources. - """ - ... + def exception(self) -> Optional[Exception]: ... + @abstractmethod + def close(self) -> None: ... -class RxSessionBase(ABC): - def __init__(self, receiver: ReceiverBase, args: SessionArgsBase): - """ - Initializes the reception session. - :param receiver: The receiver instance responsible for receiving data. - """ - self._receiver = receiver - self._base_args = args - @property - def disagg_request_id(self) -> int: - return self._base_args.params.disagg_request_id +class TxSessionBase(_SessionBase): + def __init__(self, sender: SenderBase, args: SessionArgsBase): + super().__init__(args) + self._sender = sender @abstractmethod - def receive(self, slice: KVSlice) -> concurrent.futures.Future: - """ - Receives a slice of KV cache data and returns a Future for the transfer. - :param slice: The KV slice to receive. - """ - ... + def send(self, slice: KVSlice) -> concurrent.futures.Future: ... + + +class RxSessionBase(_SessionBase): + def __init__(self, receiver: ReceiverBase, args: SessionArgsBase): + super().__init__(args) + self._receiver = receiver - @property @abstractmethod - def exception(self) -> Optional[Exception]: - """Returns any exception that occurred during the session.""" - ... + def receive(self, slice: KVSlice) -> concurrent.futures.Future: ... @abstractmethod - def close(self) -> None: - """ - Closes the session and releases any resources. - """ - ... + def wait_complete(self, blocking: bool = False) -> Optional[WaitResult]: ... diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index 0616b881120f..31c9d63cc61e 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -51,7 +51,7 @@ from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest from tensorrt_llm._torch.pyexecutor.resource_manager import KVCacheManager from tensorrt_llm._utils import nvtx_range -from tensorrt_llm.disaggregated_params import DisaggregatedParams +from tensorrt_llm.disaggregated_params import DisaggregatedParams, DisaggScheduleStyle from tensorrt_llm.runtime.generation import CUASSERT AttentionTypeCpp = tensorrt_llm.bindings.internal.batch_manager.AttentionType @@ -742,8 +742,11 @@ def __init__( params: DisaggregatedParams, sender: Sender, aux_buffer: Optional[AuxBuffer] = None, + timeout_s: Optional[float] = None, ): super().__init__(sender, SessionArgsBase(params)) + self._timeout_s = timeout_s + self._need_aux = params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST self._sender: Sender # narrow base class type for Pylance self.request_id = request_id self._aux_buffer = aux_buffer @@ -799,10 +802,10 @@ def pack_aux(self, request: LlmRequest) -> None: assert self.aux_slot is not None, "No aux_slot set for this session" self._aux_buffer.fill_slot(self.aux_slot, request) - def is_completed(self, need_aux: bool) -> bool: + def is_completed(self) -> bool: """Non-blocking check: has the transfer completed successfully?""" status = self.status - if need_aux: + if self._need_aux: return status == SessionStatus.FULLY_TRANSFERRED return status in (SessionStatus.KV_TRANSFERRED, SessionStatus.FULLY_TRANSFERRED) @@ -810,18 +813,18 @@ def has_failed(self) -> bool: """Non-blocking check: has the transfer failed?""" return self.status == SessionStatus.ERROR - def wait_complete(self, need_aux: bool, timeout: float) -> WaitResult: + def wait_complete(self) -> Optional[WaitResult]: """Block until KV (and optionally aux) transfer finishes. Returns WaitResult.COMPLETED, WaitResult.FAILED, or WaitResult.TIMEOUT. """ try: for task in self.kv_tasks: - kv_status = task.future.result(timeout=timeout) + kv_status = task.future.result(timeout=self._timeout_s) if kv_status != AgentResult.SUCCESS: return WaitResult.FAILED - if need_aux and self.aux_task is not None: - aux_status = self.aux_task.future.result(timeout=timeout) + if self._need_aux and self.aux_task is not None: + aux_status = self.aux_task.future.result(timeout=self._timeout_s) if aux_status != AgentResult.SUCCESS: return WaitResult.FAILED return WaitResult.COMPLETED @@ -1125,8 +1128,11 @@ def __init__( params: DisaggregatedParams, receiver: Receiver, aux_buffer: Optional[AuxBuffer] = None, + timeout_s: Optional[float] = None, ): super().__init__(receiver, SessionArgsBase(params)) + self._timeout_s = timeout_s + self._need_aux = params.schedule_style == DisaggScheduleStyle.GENERATION_FIRST self._receiver: Receiver # narrow base class type for Pylance self.request_id = request_id self._aux_buffer = aux_buffer @@ -1236,10 +1242,10 @@ def unpack_aux(self, request: LlmRequest) -> None: request.py_first_gen_tokens = first_gen_tokens # type: ignore[attr-defined] request.py_draft_tokens = draft_tokens # type: ignore[attr-defined] - def is_completed(self, need_aux: bool) -> bool: + def is_completed(self) -> bool: """Non-blocking check: has the transfer completed successfully?""" status = self.status - if need_aux: + if self._need_aux: return status == SessionStatus.FULLY_TRANSFERRED return status in (SessionStatus.KV_TRANSFERRED, SessionStatus.FULLY_TRANSFERRED) @@ -1247,12 +1253,12 @@ def has_failed(self) -> bool: """Non-blocking check: has the transfer failed?""" return self.status == SessionStatus.ERROR - def wait_complete(self, need_aux: bool, block_for_aux: bool = False) -> Optional[WaitResult]: - """Block until KV transfer is done; optionally wait for aux too. + def wait_complete(self, blocking: bool = False) -> Optional[WaitResult]: + """Block until transfer completes. - With block_for_aux=False (default): returns None if KV is done but aux - is still in flight — caller should re-poll next cycle. - With block_for_aux=True: spins until aux also arrives (use for block_all). + With blocking=False (default): returns None if KV is done but transfer + not fully complete — caller should re-poll next cycle. + With blocking=True: spins until fully complete. Returns WaitResult.COMPLETED on full success, WaitResult.FAILED on error. """ try: @@ -1260,17 +1266,19 @@ def wait_complete(self, need_aux: bool, block_for_aux: bool = False) -> Optional kv_status = task.future.result() if kv_status != AgentResult.SUCCESS: return WaitResult.FAILED - if need_aux: + if self._need_aux: while True: status = self.status if status == SessionStatus.FULLY_TRANSFERRED: return WaitResult.COMPLETED elif status == SessionStatus.ERROR: return WaitResult.FAILED - if not block_for_aux: + if not blocking: return None # KV done, aux still in flight; re-poll next cycle time.sleep(0.001) return WaitResult.COMPLETED + except TimeoutError: + return WaitResult.FAILED except Exception: return WaitResult.FAILED @@ -1401,10 +1409,13 @@ class TransferWorkerConfig: instance_name: str max_concurrent_sessions: int = 0 max_draft_len: Optional[int] = None + tx_timeout_s: Optional[float] = None + rx_timeout_s: Optional[float] = None class TransferWorker: def __init__(self, config: TransferWorkerConfig): + self._config = config kvm = config.kv_cache_manager self._aux_buffer = _make_aux_buffer( kvm, config.max_concurrent_sessions, config.max_draft_len @@ -1431,6 +1442,7 @@ def create_tx_session(self, request: LlmRequest) -> TxSession: params=params, sender=self._sender, aux_buffer=self._aux_buffer, + timeout_s=self._config.tx_timeout_s, ) def create_rx_session(self, request: LlmRequest) -> RxSession: @@ -1441,6 +1453,7 @@ def create_rx_session(self, request: LlmRequest) -> RxSession: params=params, receiver=self._receiver, aux_buffer=self._aux_buffer, + timeout_s=self._config.rx_timeout_s, ) def has_all_peer_req_infos_for_send(self, unique_rid: int) -> bool: diff --git a/tensorrt_llm/_torch/disaggregation/transceiver.py b/tensorrt_llm/_torch/disaggregation/transceiver.py index bcf3f2e75c09..d0df89bbac39 100644 --- a/tensorrt_llm/_torch/disaggregation/transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/transceiver.py @@ -6,7 +6,13 @@ import torch from tensorrt_llm import logger -from tensorrt_llm._torch.disaggregation.base.transfer import KVSlice, WaitResult, get_unique_rid +from tensorrt_llm._torch.disaggregation.base.transfer import ( + KVSlice, + RxSessionBase, + TxSessionBase, + WaitResult, + get_unique_rid, +) from tensorrt_llm._torch.disaggregation.native.transfer import TransferWorker, TransferWorkerConfig from tensorrt_llm._torch.disaggregation.resource.utils import get_global_layer_ids from tensorrt_llm._torch.distributed.communicator import Distributed @@ -62,6 +68,8 @@ def __init__( # can be in-flight simultaneously. AuxBuffer holds only small CPU metadata, so a # large multiplier is cheap. max_concurrent_sessions=max(1, int(kv_cache_manager.max_batch_size)) * 20000, + tx_timeout_s=self._sender_future_timeout_ms / 1000.0, + rx_timeout_s=self.kv_transfer_timeout_ms / 1000.0, ) ) self._dp_rank = mapping.tp_rank if mapping.enable_attention_dp else 0 @@ -69,8 +77,8 @@ def __init__( self._init_sync_policy() self._exchange_rank_info() - self._send_sessions = {} - self._recv_sessions = {} + self._send_sessions: Dict[int, TxSessionBase] = {} + self._recv_sessions: Dict[int, RxSessionBase] = {} self._send_reqs = {} self._recv_reqs = {} self._wait_reqs = {} @@ -190,7 +198,7 @@ def _collect_done(self, sessions: dict, reqs: dict): """Scan sessions and return (completed_rids, failed_rids).""" completed, failed = [], [] for rid, session in sessions.items(): - if session.is_completed(self._need_aux_transfer(reqs[rid])): + if session.is_completed(): completed.append(rid) elif session.has_failed(): failed.append(rid) @@ -287,12 +295,9 @@ def check_context_transfer_status( ) completed, timed_out, failed = [], [], [] - timeout = self._sender_future_timeout_ms / 1000.0 for rid in to_process: session = self._send_sessions[rid] - result = session.wait_complete( - need_aux=self._need_aux_transfer(self._send_reqs[rid]), timeout=timeout - ) + result = session.wait_complete() if result == WaitResult.COMPLETED: completed.append(rid) elif result == WaitResult.TIMEOUT: @@ -328,9 +333,7 @@ def check_gen_transfer_status(self, at_least_request_num: Optional[int]): completed, failed = [], [] for rid in to_process: - result = self._recv_sessions[rid].wait_complete( - need_aux=self._need_aux_transfer(self._recv_reqs[rid]), block_for_aux=block_all - ) + result = self._recv_sessions[rid].wait_complete(blocking=block_all) if result == WaitResult.COMPLETED: completed.append(rid) elif result == WaitResult.FAILED: