From 30113e441ba877fa99aa444871c5f14cfa3a4a40 Mon Sep 17 00:00:00 2001 From: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> Date: Fri, 20 Mar 2026 02:10:27 +0000 Subject: [PATCH 1/2] refactor the transfer.py Signed-off-by: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> --- .../_torch/disaggregation/base/transfer.py | 152 +- .../_torch/disaggregation/native/messenger.py | 1 + .../_torch/disaggregation/native/peer.py | 36 +- .../native/py_cache_transceiver.py | 101 +- .../_torch/disaggregation/native/transfer.py | 1979 ++++++++--------- .../disaggregated/test_kv_transfer.py | 63 +- .../disaggregated/test_kv_transfer_mp.py | 12 +- 7 files changed, 1148 insertions(+), 1196 deletions(-) diff --git a/tensorrt_llm/_torch/disaggregation/base/transfer.py b/tensorrt_llm/_torch/disaggregation/base/transfer.py index 3c48e83141eb..d7e2fc071093 100644 --- a/tensorrt_llm/_torch/disaggregation/base/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/base/transfer.py @@ -1,5 +1,6 @@ from __future__ import annotations +import concurrent.futures from abc import ABC, abstractmethod from dataclasses import dataclass, field from enum import Enum @@ -54,49 +55,37 @@ class KVSlice: class SessionStatus(Enum): """Status of a transfer session. - Represents the various stages/statuses that a file transfer session can go through: + Represents the lifecycle stages of a KV cache transfer session: - - INIT: The session has been initialized but not yet ready. - - READY: The session is ready to start transferring. - - TRANSFERRING: The session is in progress, currently transferring data. - - TRANSFERRED: The primary transfer has completed successfully. - - AUX_TRANSFERRED: The auxiliary part (such as tokens) of the transfer has completed successfully. - - COMPLETED: The entire session process, including all transfers, has been successfully completed. - - CANCELED: The session has been canceled by the user or system. - - ERROR: An error occurred during the session. The session could not complete successfully. + - INIT: Session initialized; waiting for the remote peer to become ready. + - READY: Peer is ready; transfer can begin. + - TRANSFERRING: KV cache transfer is in progress. + - KV_TRANSFERRED: KV cache transfer completed; auxiliary data transfer may still be pending. + - FULLY_TRANSFERRED: Both KV cache and auxiliary data (e.g. tokens) transferred successfully. + - ERROR: A transfer error occurred; the session cannot complete. """ INIT = "INIT" READY = "READY" TRANSFERRING = "TRANSFERRING" - TRANSFERRED = "TRANSFERRED" - AUX_TRANSFERRED = "AUX_TRANSFERRED" - COMPLETED = "COMPLETED" - CANCELED = "CANCELED" + KV_TRANSFERRED = "KV_TRANSFERRED" + FULLY_TRANSFERRED = "FULLY_TRANSFERRED" ERROR = "ERROR" -TaskIdType = int - - -@dataclass -class SessionState: - """State of a transfer session.""" - - status: SessionStatus - finished_tasks: List[TaskIdType] - - -class SenderBase(ABC): - """Base class for sending KV cache data.""" +class WaitResult(Enum): + """Result of waiting for a transfer session to complete.""" - ... + COMPLETED = "COMPLETED" + FAILED = "FAILED" + TIMEOUT = "TIMEOUT" -class ReceiverBase(ABC): - """Base class for receiving KV cache data.""" +@dataclass +class SessionArgsBase: + """Base arguments for transfer sessions.""" - ... + params: DisaggregatedParams def get_unique_rid(request: LlmRequest) -> Optional[int]: @@ -107,95 +96,86 @@ def get_unique_rid(request: LlmRequest) -> Optional[int]: ) -class SessionBase(ABC): - def __init__(self, request: LlmRequest): - self._request = request - self._unique_rid: Optional[int] = get_unique_rid(request) - self._state = SessionState(status=SessionStatus.INIT, finished_tasks=[]) - self._exception: Optional[Exception] = None +class SenderBase(ABC): + """Base class for sending KV cache data.""" - @property - def unique_rid(self) -> Optional[int]: - # readonly - return self._unique_rid + ... - @property - def disagg_params(self) -> Optional[DisaggregatedParams]: - return self._request.py_disaggregated_params if self._request else None - @property - def request(self) -> Optional[LlmRequest]: - return self._request +class ReceiverBase(ABC): + """Base class for receiving KV cache data.""" - @property - def state(self) -> SessionState: - """ - Returns the current state of the session. - """ - return self._state + ... - @state.setter - def state(self, state: SessionState): - """ - Set the state of the session. - :param state: The state to set. - """ - self._state = state - @abstractmethod - def poll_task(self, task_id: TaskIdType) -> SessionStatus: +class TxSessionBase(ABC): + def __init__(self, sender: SenderBase, args: SessionArgsBase): """ - Polls the status of a specific task by its ID. - :param task_id: The task ID to poll. + Initializes the transmission session. + :param sender: The sender instance responsible for sending data. + :param args: The session arguments. """ - ... + self._sender = sender + self._base_args = args + + @property + def disagg_request_id(self) -> int: + return self._base_args.params.disagg_request_id @abstractmethod - def close(self) -> None: + def send(self, slice: KVSlice) -> concurrent.futures.Future: """ - Closes the session and releases any resources. + Sends a slice of KV cache data and returns a Future for the transfer. + :param slice: The KV slice to send. """ ... @property + @abstractmethod def exception(self) -> Optional[Exception]: """ Returns any exception that occurred during the session. """ - return self._exception - - -class TxSessionBase(SessionBase): - def __init__(self, sender: SenderBase, request: LlmRequest): - """ - Initializes the transmission session. - :param sender: The sender instance responsible for sending data. - :param request: The LLM request associated with this session. - """ - self._sender = sender - super().__init__(request) + ... @abstractmethod - def send(self, slice: KVSlice) -> TaskIdType: + def close(self) -> None: """ - Sends a slice of KV cache data and returns the task ID. - :param slice: The KV slice to send. + Closes the session and releases any resources. """ + ... -class RxSessionBase(SessionBase): - def __init__(self, receiver: ReceiverBase, request: LlmRequest): +class RxSessionBase(ABC): + def __init__(self, receiver: ReceiverBase, args: SessionArgsBase): """ Initializes the reception session. :param receiver: The receiver instance responsible for receiving data. """ - super().__init__(request) self._receiver = receiver + self._base_args = args + + @property + def disagg_request_id(self) -> int: + return self._base_args.params.disagg_request_id @abstractmethod - def receive(self, slice: KVSlice) -> TaskIdType: + def receive(self, slice: KVSlice) -> concurrent.futures.Future: """ - Receives a slice of KV cache data and returns the task ID. + Receives a slice of KV cache data and returns a Future for the transfer. :param slice: The KV slice to receive. """ ... + + @property + @abstractmethod + def exception(self) -> Optional[Exception]: + """Returns any exception that occurred during the session.""" + ... + + @abstractmethod + def close(self) -> None: + """ + Closes the session and releases any resources. + """ + ... diff --git a/tensorrt_llm/_torch/disaggregation/native/messenger.py b/tensorrt_llm/_torch/disaggregation/native/messenger.py index da2185d5259f..ceb6aa626ed9 100644 --- a/tensorrt_llm/_torch/disaggregation/native/messenger.py +++ b/tensorrt_llm/_torch/disaggregation/native/messenger.py @@ -178,6 +178,7 @@ def stop(self, timeout: int = 5) -> None: def _close_socket(socket: zmq.Socket) -> None: try: if not socket.closed: + socket.setsockopt(zmq.LINGER, 0) socket.close() except Exception as e: logger.error(f"Error closing socket: {e}") diff --git a/tensorrt_llm/_torch/disaggregation/native/peer.py b/tensorrt_llm/_torch/disaggregation/native/peer.py index 1cc015ae1f78..a021a3f113de 100644 --- a/tensorrt_llm/_torch/disaggregation/native/peer.py +++ b/tensorrt_llm/_torch/disaggregation/native/peer.py @@ -120,12 +120,15 @@ def get_pool_mapping(self, peer_ri: RankInfo) -> Dict[LGPoolKey, LGPoolKey]: self_pt = self._self_ext_cache.page_table peer_pt = peer_ri.page_table + if self_pt is None or peer_pt is None: + self._lg_pool_mapping_cache[key] = mapping + return mapping if not self_pt.layer_groups or not peer_pt.layer_groups: - mapping[(0, 0)] = (0, 0) self._lg_pool_mapping_cache[key] = mapping return mapping peer_layer_to_group = get_layer_to_layer_group(peer_pt) + assert self._ri.attention is not None kv_factor = self._ri.attention.kv_factor for self_lg_idx, self_lg in enumerate(self_pt.layer_groups): @@ -199,6 +202,8 @@ def get_kv_map( self_pt = self._self_ext_cache.page_table peer_pt = peer_ri.page_table + assert self_pt is not None + assert peer_pt is not None self_lg_idx, self_pi = self_pool_key peer_lg_idx, peer_pi = peer_pool_key self_lg = self_pt.layer_groups[self_lg_idx] @@ -206,6 +211,7 @@ def get_kv_map( self_pv = self_lg.pool_views[self_pi] peer_pv = peer_lg.pool_views[peer_pi] + assert self._ri.attention is not None kv_factor = self._ri.attention.kv_factor is_indexer = len(self_pv.buffer_entries) == 0 self_pool_role = ( @@ -324,3 +330,31 @@ def get_peer_overlap(self, peer_rank_info: RankInfo, peer_dp_rank: int) -> PeerO ) self._overlap_cache[key] = targets return targets + + def should_send_kv(self, peer_overlap: PeerOverlap, peer_rank_info: RankInfo) -> bool: + dup_head_factor = peer_overlap.duplicate_head_factor + if dup_head_factor <= 1: + return True + self_tp_rank_in_dp_group = self._ri.tp_rank % self._ri.tp_size_per_dp_group + return (peer_rank_info.dp_rank % dup_head_factor) == ( + self_tp_rank_in_dp_group % dup_head_factor + ) + + def should_send_aux(self, peer_rank_info: RankInfo) -> bool: + # to ensure the transfer aux is not duplicated + + # TP: only the first rank in each peer-TP-sized group sends aux + ratio = max(1, self._ri.tp_size_per_dp_group // peer_rank_info.tp_size_per_dp_group) + self_tp_rank_in_dp_group = self._ri.tp_rank % self._ri.tp_size_per_dp_group + should_send_in_tp = self_tp_rank_in_dp_group % ratio == 0 + + # PP: only the first self-PP rank whose layers overlap with the peer's PP rank sends aux. + # All tp/pp ranks have the same aux data, so pick the first overlapping one to avoid duplication. + peer_start_layer = sum(peer_rank_info.layer_num_per_pp[: peer_rank_info.pp_rank]) + peer_end_layer = peer_start_layer + peer_rank_info.layer_num_per_pp[peer_rank_info.pp_rank] + offset = 0 + for p, n in enumerate(self._ri.layer_num_per_pp): + if offset < peer_end_layer and offset + n > peer_start_layer: + return should_send_in_tp and p == self._ri.pp_rank + offset += n + return False diff --git a/tensorrt_llm/_torch/disaggregation/native/py_cache_transceiver.py b/tensorrt_llm/_torch/disaggregation/native/py_cache_transceiver.py index 5fb0bcdee3ae..4b2570d83bd0 100644 --- a/tensorrt_llm/_torch/disaggregation/native/py_cache_transceiver.py +++ b/tensorrt_llm/_torch/disaggregation/native/py_cache_transceiver.py @@ -1,4 +1,3 @@ -import concurrent import uuid from collections import defaultdict from itertools import chain @@ -8,7 +7,7 @@ import tensorrt_llm from tensorrt_llm import logger -from tensorrt_llm._torch.disaggregation.base.transfer import KVSlice, SessionStatus, get_unique_rid +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 @@ -92,7 +91,7 @@ def __init__( 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.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) @@ -107,7 +106,7 @@ def __init__( 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 + 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) @@ -120,13 +119,11 @@ def __init__( 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.send_task_ids = {} # request_id to send_task_id self.recv_sessions = {} # request_id to recv_session - self.recv_task_ids = {} # request_id to recv_task_id 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._rank_info.page_table + 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") @@ -195,10 +192,10 @@ def respond_and_send_async(self, req: LlmRequest): # stores block ids, not raw data kv_slice = self._create_kv_slice(req) # sending actual kv data - send_task_id = send_session.send(kv_slice) + send_session.send(kv_slice) if self._need_aux_transfer(req): + send_session.pack_aux(req) send_session.send_aux() - self.send_task_ids[unique_rid] = send_task_id # contains metadata about itself so the gen server can see req.context_phase_params = ContextPhaseParams( @@ -225,8 +222,7 @@ def request_and_receive_async(self, req: LlmRequest): recv_session = self.transfer_worker.create_rx_session(req) self.recv_sessions[unique_rid] = recv_session kv_slice = self._create_kv_slice(req) - recv_task_id = recv_session.receive(kv_slice) - self.recv_task_ids[unique_rid] = recv_task_id + 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): @@ -239,15 +235,9 @@ def check_context_transfer_status(self, at_least_request_num: int, mark_complete 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) - session_status = session.state.status - if need_aux: - if session_status == SessionStatus.AUX_TRANSFERRED: - local_completed_request_ids.append(request_id) - elif session_status == SessionStatus.ERROR: - local_failed_request_ids.append(request_id) - elif session_status == SessionStatus.TRANSFERRED: + if session.is_completed(need_aux): local_completed_request_ids.append(request_id) - elif session_status == SessionStatus.ERROR: + elif session.has_failed(): local_failed_request_ids.append(request_id) local_sync_request_ids = local_completed_request_ids + local_failed_request_ids @@ -269,20 +259,20 @@ def check_context_transfer_status(self, at_least_request_num: int, mark_complete 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] - try: - if session.wait_complete( - self.send_task_ids[request_id], - wait_aux=True, - timeout_ms=self.sender_future_timeout_ms, - ): - completed_request_ids.append(request_id) - except concurrent.futures.TimeoutError: - logger.warning(f"TxSession {session.unique_rid} timed out waiting for completion") + 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) - except Exception: - logger.warning(f"TxSession {session.unique_rid} failed to complete") + 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: @@ -296,7 +286,6 @@ def check_context_transfer_status(self, at_least_request_num: int, mark_complete 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] - del self.send_task_ids[request_id] return completed_request_ids, failed_request_ids @@ -310,15 +299,9 @@ def check_gen_transfer_status(self, at_least_request_num: int): 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) - session_status = session.state.status - if need_aux: - if session_status == SessionStatus.AUX_TRANSFERRED: - local_completed_request_ids.append(request_id) - elif session_status == SessionStatus.ERROR: - local_failed_request_ids.append(request_id) - elif session_status == SessionStatus.TRANSFERRED: + if session.is_completed(need_aux): local_completed_request_ids.append(request_id) - elif session_status == SessionStatus.ERROR: + elif session.has_failed(): local_failed_request_ids.append(request_id) local_sync_request_ids = local_completed_request_ids + local_failed_request_ids @@ -350,25 +333,43 @@ def check_gen_transfer_status(self, at_least_request_num: int): completed_request_ids = [] failed_request_ids = [] for request_id in to_complete_request_ids: - recv_task_id = self.recv_task_ids[request_id] recv_session = self.recv_sessions[request_id] req = self.recv_req_id_to_request[request_id] - if recv_session.wait_complete(recv_task_id, wait_aux=self._need_aux_transfer(req)): + 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) - else: + 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: - self.recv_req_id_to_request[ - request_id - ].state = LlmRequestState.DISAGG_GENERATION_TRANS_COMPLETE + 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: - self.recv_req_id_to_request[request_id].state = LlmRequestState.DISAGG_TRANS_ERROR + req.state = LlmRequestState.DISAGG_TRANS_ERROR del self.recv_req_id_to_request[request_id] - self.transfer_worker.clear_session(self.recv_sessions[request_id]) + self.transfer_worker.clear_session(recv_session) del self.recv_sessions[request_id] - del self.recv_task_ids[request_id] return @@ -384,7 +385,9 @@ def get_disaggregated_params(self) -> Dict[str, Any]: # requests before context-phase response data arrives. return { "ctx_dp_rank": self.dp_rank, - "ctx_info_endpoint": self.context_info_endpoint, + "ctx_info_endpoint": [self.context_info_endpoint] + if self.context_info_endpoint + else None, } def prepare_context_requests(self, requests: List[LlmRequest]): diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index 01a201bd0d32..f005509cce0b 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -1,17 +1,18 @@ -import concurrent +from __future__ import annotations + +import concurrent.futures import os import queue import threading +import time import weakref -from dataclasses import asdict, dataclass, field +from dataclasses import asdict, dataclass from enum import Enum from typing import List, Optional import msgpack import torch -from tensorrt_llm.bindings.executor import ContextPhaseParams - try: from cuda.bindings import runtime as cudart except ImportError: @@ -29,15 +30,18 @@ ) from tensorrt_llm._torch.disaggregation.base.transfer import ( KVSlice, + ReceiverBase, RxSessionBase, + SenderBase, + SessionArgsBase, SessionStatus, - TaskIdType, TxSessionBase, + WaitResult, ) -from tensorrt_llm._torch.disaggregation.native.auxiliary import AuxBuffer, AuxSlot +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 PeerOverlap, PeerRegistrar +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 @@ -50,14 +54,13 @@ 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.disaggregated_params import DisaggregatedParams, DisaggScheduleStyle +from tensorrt_llm.disaggregated_params import DisaggregatedParams from tensorrt_llm.runtime.generation import CUASSERT AttentionTypeCpp = tensorrt_llm.bindings.internal.batch_manager.AttentionType LlmRequestType = tensorrt_llm.bindings.internal.batch_manager.LlmRequestType -# Environment variable to control the number of threads for task queue processing -# Default is 1 (single-threaded, original behavior) +# Number of worker threads for KV transfer queues (default: 1) KV_TRANSFER_NUM_THREADS = int(os.environ.get("TRTLLM_KV_TRANSFER_NUM_THREADS", "1")) @@ -66,7 +69,7 @@ class RecvReqInfo: sender_req_id: int instance_name: str instance_rank: int - block_ids_per_layer_groups: list[list[int]] # Block IDs per layer group + block_ids_per_layer_groups: list[list[int]] unique_rid: int start_token_idx: Optional[int] = None aux_slot: Optional[int] = None @@ -83,411 +86,95 @@ def from_bytes(cls, data: bytes) -> "RecvReqInfo": class ReadMeta: unique_rid: int slice_id: int - target_ranks: List[int] = None + target_ranks: Optional[List[int]] = None + + +class WriteMetaType(Enum): + KV = "KV" + AUX = "AUX" @dataclass class WriteMeta: - future_for_task: concurrent.futures.Future - + task_future: concurrent.futures.Future expected_transfers: int peer_name: str peer_rank: int - src_kv_ptrs: List[int] = None - dst_kv_ptrs: List[int] = None - kv_sizes: List[int] = None - dst_device_id: int = None - src_aux_ptrs: List[int] = None - dst_aux_ptrs: List[int] = None - aux_sizes: List[int] = None - peer_endpoint: Optional[str] = None # used for send state - unique_rid: Optional[int] = None + peer_endpoint: str + unique_rid: int + src_ptrs: List[int] + dst_ptrs: List[int] + sizes: List[int] + dst_device_id: Optional[int] = None slice_id: Optional[int] = None - is_last_slice: Optional[bool] = False - is_only_aux: Optional[bool] = False + is_last_slice: bool = False + meta_type: WriteMetaType = WriteMetaType.KV class MessageType: TERMINATION = b"TERMINATION" - TASK_STATUS = b"TASK_STATUS" - PEER_INFO = b"PEER_INFO" + KV_AGENT_RESULT = b"KV_AGENT_RESULT" REQUEST_DATA = b"REQUEST_DATA" REQUEST_INSTANCE_INFO = b"REQUEST_INSTANCE_INFO" REGISTER_RANK_INFO = b"REGISTER_RANK_INFO" - AUX_SEND_STATUS = b"AUX_SEND_STATUS" + AUX_AGENT_RESULT = b"AUX_AGENT_RESULT" class TaskStatus(Enum): INIT = "INIT" TRANSFERRING = "TRANSFERRING" TRANSFERRED = "TRANSFERRED" - AUX_TRANSFERRED = "AUX_TRANSFERRED" - COMPLETED = "COMPLETED" - CANCELED = "CANCELED" ERROR = "ERROR" -class AuxSendTask: - def __init__(self, unique_rid: int, slot_id: int, peer_registrar: PeerRegistrar): - self._unique_rid = unique_rid - self._slot_id = slot_id - self._registrar = peer_registrar - self._status = TaskStatus.INIT - self._future = concurrent.futures.Future() - self._expected_transfers = None - self._transferred_count = 0 - self._perf_timer = PerfTimer() if perf_log_manager.enabled else None - - def is_active(self) -> bool: - return ( - self._expected_transfers is None or self._transferred_count < self._expected_transfers - ) - - @property - def status(self) -> TaskStatus: - return self._status - - @status.setter - def status(self, s: TaskStatus): - self._status = s - - @property - def future(self) -> concurrent.futures.Future: - return self._future - - def _create_write_meta(self, req_info: RecvReqInfo) -> WriteMeta: - peer_rank_info = self._registrar.get_peer_rank_info( - req_info.instance_name, req_info.instance_rank - ) - if self._perf_timer is not None: - self._perf_timer.record_prepare_args_start(peer_rank_info.instance_rank) - expected_transfers = len( - self._registrar.get_peer_overlap(peer_rank_info, peer_rank_info.dp_rank).ranks - ) - if self._expected_transfers is None: - self._expected_transfers = expected_transfers - if not self._should_write(peer_rank_info): - if self._perf_timer is not None: - self._perf_timer.record_prepare_args_end(peer_rank_info.instance_rank) - self._perf_timer.record_transfer_sizes( - req_info.instance_name + str(req_info.instance_rank), 0, 0 - ) - return WriteMeta( - future_for_task=self._future, - src_aux_ptrs=[], - dst_aux_ptrs=[], - aux_sizes=[], - expected_transfers=expected_transfers, - is_only_aux=True, - peer_name=req_info.instance_name + str(req_info.instance_rank), - peer_rank=req_info.instance_rank, - peer_endpoint=self._registrar.get_peer_rank_info( - req_info.instance_name, req_info.instance_rank - ).self_endpoint, - unique_rid=self._unique_rid, - ) - peer_aux_meta = self._registrar.get_peer_rank_info( - req_info.instance_name, req_info.instance_rank - ).aux_meta - - peer_slot = req_info.aux_slot - - src_aux_meta = self._registrar.self_rank_info.aux_meta - - src_ptrs = [ - ptr + item_size * self._slot_id - for ptr, item_size in zip(src_aux_meta.ptrs, src_aux_meta.item_sizes) - ] - dst_ptrs = [ - ptr + item_size * peer_slot - for ptr, item_size in zip(peer_aux_meta.ptrs, peer_aux_meta.item_sizes) - ] - size = [item_size for item_size in src_aux_meta.item_sizes] - - if self._perf_timer is not None: - self._perf_timer.record_prepare_args_end(peer_rank_info.instance_rank) - self._perf_timer.record_transfer_sizes(req_info.instance_rank, sum(size), len(src_ptrs)) - return WriteMeta( - future_for_task=self._future, - src_aux_ptrs=src_ptrs, - dst_aux_ptrs=dst_ptrs, - aux_sizes=size, - expected_transfers=expected_transfers, - is_only_aux=True, - peer_name=req_info.instance_name + str(req_info.instance_rank), - peer_rank=req_info.instance_rank, - peer_endpoint=self._registrar.get_peer_rank_info( - req_info.instance_name, req_info.instance_rank - ).self_endpoint, - unique_rid=self._unique_rid, - ) - - def _should_write(self, peer_rank_info: RankInfo) -> bool: - # to ensure the transfer aux is not duplicated - self_ri = self._registrar.self_rank_info - self_tp_rank_in_dp_group = self_ri.tp_rank % self_ri.tp_size_per_dp_group +class AgentResult(Enum): + SUCCESS = "SUCCESS" + FAILED = "FAILED" - should_send_in_tp = False - if self_ri.tp_size_per_dp_group <= peer_rank_info.tp_size_per_dp_group: - should_send_in_tp = True - - else: - ratio = self_ri.tp_size_per_dp_group // peer_rank_info.tp_size_per_dp_group - should_send_in_tp = self_tp_rank_in_dp_group % ratio == 0 - - # Compute peer pp_rank's global layer range from layer_num_per_pp - self_layer_num_per_pp = self_ri.layer_num_per_pp - peer_layer_num_per_pp = peer_rank_info.layer_num_per_pp - peer_start_layer = sum(peer_layer_num_per_pp[: peer_rank_info.pp_rank]) - peer_end_layer = peer_start_layer + peer_layer_num_per_pp[peer_rank_info.pp_rank] - - # Find the first self pp_rank whose global layers overlap with peer's pp_rank. - # All tp/pp ranks have the same aux data, so we only select one self pp_rank - # to send to the peer; pick the first overlapping one to avoid duplication. - first_matching_self_pp_rank = None - self_layer_offset = 0 - for p in range(self_ri.pp_size): - self_start = self_layer_offset - self_end = self_start + self_layer_num_per_pp[p] - if self_start < peer_end_layer and self_end > peer_start_layer: - first_matching_self_pp_rank = p - break - self_layer_offset += self_layer_num_per_pp[p] - should_send_in_pp = first_matching_self_pp_rank == self_ri.pp_rank - return should_send_in_tp and should_send_in_pp +class SendTaskBase: + def __init__(self, params: DisaggregatedParams): + self.status = TaskStatus.INIT + self.future = concurrent.futures.Future() + self._params = params + assert params.disagg_request_id is not None + self._unique_rid: int = params.disagg_request_id + self._perf_timer = PerfTimer() if perf_log_manager.enabled else None - def print_perf_info(self, peer_rank: int): - ri = self._registrar.self_rank_info + def print_perf_info(self, peer_rank: int, instance_name: str, instance_rank: int): + if self._perf_timer is None: + return perf_log_manager.log_task_perf( - "AuxSendTask", + type(self).__name__, self._unique_rid, peer_rank, - ri.instance_name, - ri.instance_rank, + instance_name, + instance_rank, self._perf_timer, ) -class KVSendTask: +class AuxSendTask(SendTaskBase): + def __init__(self, params: DisaggregatedParams, slot: Optional[int]): + super().__init__(params) + self._slot = slot + self._transfer_count = 0 + + +class KVSendTask(SendTaskBase): def __init__( self, kv_slice: KVSlice, - unique_rid: int, + params: DisaggregatedParams, slice_id: int, - peer_registrar: PeerRegistrar, ): - self._registrar = peer_registrar - self._future = concurrent.futures.Future() - self._first_transfer = False - self._extraction_count = 0 - self._expected_transfers = 0 + super().__init__(params) + self.slice_id = slice_id + self.transferred_count = 0 self._slice = kv_slice - self._unique_rid = unique_rid - self._slice_id = slice_id - self._status = TaskStatus.INIT - self._transferred_count = 0 - self._perf_timer = PerfTimer() if perf_log_manager.enabled else None - - @property - def status(self) -> TaskStatus: - return self._status - - @status.setter - def status(self, s: TaskStatus): - self._status = s - - @property - def future(self) -> concurrent.futures.Future: - return self._future - - @property - def slice_id(self) -> int: - return self._slice_id - - @property - def transferred_count(self) -> int: - return self._transferred_count - - @transferred_count.setter - def transferred_count(self, v: int): - self._transferred_count = v - - @nvtx_range("_create_write_meta") - def _create_write_meta(self, req_info: RecvReqInfo) -> WriteMeta: - assert self.is_active(), ( - f"KVSendTask {self._unique_rid}:{self._slice_id} is not active, first_transfer: {self._first_transfer}, " - f"extraction_count: {self._extraction_count}, expected_transfers: {self._expected_transfers}" - ) - peer_ri = self._registrar.get_peer_rank_info(req_info.instance_name, req_info.instance_rank) - if self._perf_timer is not None: - self._perf_timer.record_prepare_args_start(peer_ri.instance_rank) - targets = self._registrar.get_peer_overlap(peer_ri, peer_ri.dp_rank) - expected_transfers = len(targets.ranks) - if not self._first_transfer: - self._first_transfer = True - self._expected_transfers = expected_transfers - self._extraction_count = self._extraction_count + 1 - if not self._should_write(targets, peer_ri): - if self._perf_timer is not None: - self._perf_timer.record_prepare_args_end(peer_ri.instance_rank) - self._perf_timer.record_transfer_sizes(peer_ri.instance_rank, 0, 0) - return WriteMeta( - future_for_task=self._future, - src_kv_ptrs=[], - dst_kv_ptrs=[], - kv_sizes=[], - expected_transfers=expected_transfers, - peer_name=peer_ri.instance_name + str(peer_ri.instance_rank), - peer_rank=peer_ri.instance_rank, - peer_endpoint=peer_ri.self_endpoint, - unique_rid=self._unique_rid, - slice_id=self._slice_id, - is_last_slice=self._slice.is_last_slice, - ) - - dst_device_id = peer_ri.device_id - dst_block_ids_per_groups = req_info.block_ids_per_layer_groups - src_block_ids_per_groups = self._slice.block_ids_per_layer_groups - - extractor = self._registrar.self_extractor - 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) - - # Aggregate fragments from all matching pools - src_frags: List[int] = [] - dst_frags: List[int] = [] - kv_sizes: List[int] = [] - - 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] - - if len(src_block_ids) + 1 == len(dst_block_ids): - # FIXME: this is a temporary solution, need to be fixed for the draft tokens - logger.warning( - "src_block_num is one less than dst_block_num, maybe it is due to draft tokens," - " remove the last block from dst_block_ids " - ) - dst_block_ids = dst_block_ids[:-1] - src_block_ids, dst_block_ids = self._filter_kv_blocks(src_block_ids, dst_block_ids) - - src_region = extractor.extract(src_block_ids, layer_group_id=self_lg, pool_idx=self_pi) - dst_region = peer_extractor.extract( - dst_block_ids, layer_group_id=peer_lg, pool_idx=peer_pi - ) - mapper = self._registrar.get_kv_map(peer_ri, (self_lg, self_pi), (peer_lg, peer_pi)) - region_pair = mapper.map(src_region, dst_region) - region_pairs = region_pair if isinstance(region_pair, list) else [region_pair] - for rp in region_pairs: - src_frags.extend(rp.src.memory.ptrs) - dst_frags.extend(rp.dst.memory.ptrs) - frag_size = rp.src.memory.bytes_per_region - kv_sizes.extend([frag_size] * len(rp.src.memory.ptrs)) - - if self._perf_timer is not None: - transfer_total_size = sum(kv_sizes) - self._perf_timer.record_prepare_args_end(peer_ri.instance_rank) - self._perf_timer.record_transfer_sizes( - peer_ri.instance_rank, transfer_total_size, len(dst_frags) - ) - - return WriteMeta( - future_for_task=self._future, - src_kv_ptrs=src_frags, - dst_kv_ptrs=dst_frags, - kv_sizes=kv_sizes, - dst_device_id=dst_device_id, - expected_transfers=expected_transfers, - peer_name=peer_ri.instance_name + str(peer_ri.instance_rank), - peer_rank=peer_ri.instance_rank, - peer_endpoint=peer_ri.self_endpoint, - unique_rid=self._unique_rid, - slice_id=self._slice_id, - is_last_slice=self._slice.is_last_slice, - ) - - def is_active(self) -> bool: - if self._first_transfer: - return self._extraction_count < self._expected_transfers - else: - return True - - def _should_write(self, peer_overlap: PeerOverlap, peer_rank_info: RankInfo) -> bool: - dup_head_factor = peer_overlap.duplicate_head_factor - if dup_head_factor <= 1: - return True - peer_ri = peer_rank_info - peer_dp_rank = peer_ri.dp_rank - self_ri = self._registrar.self_rank_info - self_tp_rank_in_dp_group = self_ri.tp_rank % self_ri.tp_size_per_dp_group - return (peer_dp_rank % dup_head_factor) == (self_tp_rank_in_dp_group % dup_head_factor) - - def _filter_kv_blocks(self, src_block_ids, dst_block_ids) -> tuple[list[int], list[int]]: - # TODO: filter the kv block_ids according to the peer_overlap - return src_block_ids, dst_block_ids - - def print_perf_info(self, peer_rank: int): - ri = self._registrar.self_rank_info - perf_log_manager.log_task_perf( - "KVSendTask", - self._unique_rid, - peer_rank, - ri.instance_name, - ri.instance_rank, - self._perf_timer, - ) - - -@dataclass -class ReqInfoManager: - # unique_rid -> instance_rank -> RecvReqInfo - _peer_requests: dict[str, dict[int, RecvReqInfo]] = field(default_factory=dict) - _lock = threading.Lock() - - def add_req_info(self, unique_rid: str, instance_rank: int, req_info: RecvReqInfo): - with self._lock: - if unique_rid not in self._peer_requests: - self._peer_requests[unique_rid] = {} - self._peer_requests[unique_rid][instance_rank] = req_info - - def is_ready(self, unique_rid: str, expected_count: int) -> bool: - with self._lock: - requests = self._peer_requests.get(unique_rid) - if not requests: - return False - return len(requests) == expected_count - def get_req_info(self, unique_rid: str) -> Optional[dict[int, RecvReqInfo]]: - with self._lock: - return self._peer_requests.get(unique_rid, {}) - def get_first_req_info(self, unique_rid: str) -> Optional[RecvReqInfo]: - with self._lock: - reqs = self._peer_requests.get(unique_rid) - if not reqs: - return None - return next(iter(reqs.values())) - - def remove_req_info(self, unique_rid: str): - with self._lock: - if unique_rid in self._peer_requests: - del self._peer_requests[unique_rid] - - -def _handle_error(exception: Exception): - import traceback - - logger.error( - f"Exception in Sender._start_listener: {exception}\nTraceback: {traceback.format_exc()}" - ) - - -class Sender: +class Sender(SenderBase): def __init__( self, peer_registrar: PeerRegistrar, @@ -497,20 +184,16 @@ def __init__( self._registrar = peer_registrar self._device_id = device_id self._agent = agent - self._peer_reqs = ReqInfoManager() - + # unique_rid -> instance_rank -> RecvReqInfo + self._peer_requests: dict = {} + self._peer_requests_lock = threading.Lock() self._messenger = ZMQMessenger(mode="ROUTER") - self._start_listener() - self._dealers = {} - self._tx_sessions = {} # unique_rid -> TxSession - self._sessions_lock = threading.Lock() # Protects _tx_sessions access - logger.info(f" Sender init end with endpoint: {self._messenger.endpoint}") - self._closed = False + self._sessions = {} # unique_rid -> TxSession + self._sessions_lock = threading.Lock() # Protects _sessions access + self._shutdown = False self._instance_rank = self._registrar.self_rank_info.instance_rank self._loaded_remote_agents: set[str] = set() - - # Multi-threaded task queue support self._num_threads = KV_TRANSFER_NUM_THREADS self._send_task_queues: List[queue.Queue] = [ queue.Queue() for _ in range(self._num_threads) @@ -519,25 +202,65 @@ def __init__( threading.Thread(target=self._process_task_queue, args=(i,), daemon=True) for i in range(self._num_threads) ] + + self._start_listener() for t in self._worker_threads: t.start() - logger.info(f"Sender started with {self._num_threads} worker thread(s)") + logger.info( + f"Sender init end with endpoint: {self._messenger.endpoint}," + f" {self._num_threads} worker thread(s)" + ) @property def endpoint(self): return self._messenger.endpoint - def setup_session(self, tx_session: TxSessionBase): - unique_rid = tx_session.unique_rid + def _add_req_info(self, unique_rid: int, instance_rank: int, req_info: RecvReqInfo): + with self._peer_requests_lock: + if unique_rid not in self._peer_requests: + self._peer_requests[unique_rid] = {} + self._peer_requests[unique_rid][instance_rank] = req_info + + def _is_req_ready(self, unique_rid: int, expected_count: int) -> bool: + with self._peer_requests_lock: + requests = self._peer_requests.get(unique_rid) + if not requests: + return False + return len(requests) == expected_count + + def _get_req_info(self, unique_rid: Optional[int]) -> Optional[dict]: + with self._peer_requests_lock: + return self._peer_requests.get(unique_rid) + + def _get_first_req_info(self, unique_rid: Optional[int]) -> Optional[RecvReqInfo]: + with self._peer_requests_lock: + reqs = self._peer_requests.get(unique_rid) + if not reqs: + return None + return next(iter(reqs.values())) + + def _remove_req_info(self, unique_rid: int): + with self._peer_requests_lock: + self._peer_requests.pop(unique_rid, None) + + def setup_session(self, tx_session: "TxSession"): + unique_rid = tx_session.disagg_request_id with self._sessions_lock: - self._tx_sessions[unique_rid] = weakref.ref(tx_session) - req_info = self._peer_reqs.get_first_req_info(unique_rid) - if req_info: - if self._has_all_peer_req_infos(req_info): - tx_session.state.status = SessionStatus.READY - - def _get_tx_session(self, unique_rid: int) -> TxSessionBase: - session_ref = self._tx_sessions.get(unique_rid) + self._sessions[unique_rid] = weakref.ref(tx_session) + + req_info = self._get_first_req_info(unique_rid) + + if req_info: + peer_ri = self._registrar.get_peer_rank_info( + req_info.instance_name, req_info.instance_rank + ) + expected_count = len(self._registrar.get_peer_overlap(peer_ri, peer_ri.dp_rank).ranks) + if self._is_req_ready(unique_rid, expected_count): + tx_session.receiver_ready = True + return + + def _get_session(self, unique_rid: Optional[int]) -> Optional["TxSession"]: + session_ref = self._sessions.get(unique_rid) if session_ref is None: return None session = session_ref() @@ -546,224 +269,336 @@ def _get_tx_session(self, unique_rid: int) -> TxSessionBase: return None return session - def submit_task(self, agent_args: WriteMeta): + def _enqueue(self, write_meta: WriteMeta): # Distribute tasks to threads by unique_rid to ensure same session's tasks # are processed by the same thread in order - thread_idx = agent_args.unique_rid % self._num_threads - self._send_task_queues[thread_idx].put(agent_args) + thread_idx = write_meta.unique_rid % self._num_threads + self._send_task_queues[thread_idx].put(write_meta) def _process_task_queue(self, thread_idx: int): - """Process tasks from the queue assigned to this thread. - - Args: - thread_idx: Index of the worker thread (0 to num_threads-1) - """ device_id = self._device_id torch.cuda.set_device(device_id) CUASSERT(cudart.cudaSetDevice(device_id)) task_queue = self._send_task_queues[thread_idx] while True: - agent_args = task_queue.get() - if agent_args is None: + write_meta = task_queue.get() + if write_meta is None: break - if agent_args.is_only_aux: - logger.debug( - f"_process_task_queue[{thread_idx}]: delivering aux task to agent: {agent_args}" + try: + if write_meta.meta_type == WriteMetaType.AUX: + logger.debug( + f"_process_task_queue[{thread_idx}]: delivering aux task to agent: {write_meta}" + ) + self._deliver_aux_to_agent(write_meta) + else: + self._deliver_kv_to_agent(write_meta) + except Exception as e: + logger.error( + f"_process_task_queue[{thread_idx}]: unhandled exception for " + f"unique_rid={write_meta.unique_rid}: {e}" ) - self._deliver_aux_to_agent(agent_args) - else: - self._deliver_kv_to_agent(agent_args) + if not write_meta.task_future.done(): + write_meta.task_future.set_exception(e) @staticmethod @nvtx_range("_make_agent_request") - def _make_agent_request(agent_args: WriteMeta, is_aux: bool, device_id: int): - if is_aux: - assert agent_args.src_aux_ptrs is not None and agent_args.dst_aux_ptrs is not None - src_list = [ - (src_ptr, size, 0) - for src_ptr, size in zip(agent_args.src_aux_ptrs, agent_args.aux_sizes) - ] - dst_list = [ - (dst_ptr, size, 0) - for dst_ptr, size in zip(agent_args.dst_aux_ptrs, agent_args.aux_sizes) - ] - src_mem_type = MemoryType.DRAM - dst_mem_type = MemoryType.DRAM - peer_name = agent_args.peer_name + def _make_agent_request(write_meta: WriteMeta, device_id: int) -> "TransferRequest": + if write_meta.meta_type == WriteMetaType.AUX: + src_dev, dst_dev, mem_type = 0, 0, MemoryType.DRAM else: - assert agent_args.src_kv_ptrs is not None and agent_args.dst_kv_ptrs is not None - src_list = [ - (src_ptr, size, device_id) - for src_ptr, size in zip(agent_args.src_kv_ptrs, agent_args.kv_sizes) - ] - dst_list = [ - (dst_ptr, size, agent_args.dst_device_id) - for dst_ptr, size in zip(agent_args.dst_kv_ptrs, agent_args.kv_sizes) - ] - src_mem_type = MemoryType.VRAM - dst_mem_type = MemoryType.VRAM - peer_name = agent_args.peer_name - - # Use C++ MemoryDescs directly with batch constructor (list of tuples) - src_memory_descs = MemoryDescs(src_mem_type, src_list) - dst_memory_descs = MemoryDescs(dst_mem_type, dst_list) - request = TransferRequest( - TransferOp.WRITE, src_memory_descs, dst_memory_descs, peer_name, None + if write_meta.dst_device_id is None: + raise RuntimeError( + f"_make_agent_request: dst_device_id is None for KV transfer " + f"unique_rid={write_meta.unique_rid}" + ) + 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) + ] + dst_list = [ + (ptr, size, dst_dev) for ptr, size in zip(write_meta.dst_ptrs, write_meta.sizes) + ] + return TransferRequest( + TransferOp.WRITE, + MemoryDescs(mem_type, src_list), + MemoryDescs(mem_type, dst_list), + write_meta.peer_name, + None, ) - return request, src_list, dst_list @nvtx_range("_deliver_kv_to_agent") - def _deliver_kv_to_agent(self, agent_args: WriteMeta): - assert len(agent_args.src_kv_ptrs) == len(agent_args.dst_kv_ptrs) - assert len(agent_args.kv_sizes) == len(agent_args.src_kv_ptrs) - assert agent_args.is_only_aux is False, "agent_args.is_only_aux should be False" - - unique_rid = agent_args.unique_rid - slice_id = agent_args.slice_id - peer_endpoint = agent_args.peer_endpoint - session = self._get_tx_session(unique_rid) - assert session is not None - task = session._kv_tasks[slice_id] - if task._perf_timer is not None: - task._perf_timer.record_push_end(agent_args.peer_rank) - assert session.state.status != SessionStatus.ERROR - session.state.status = SessionStatus.TRANSFERRING - task.status = TaskStatus.TRANSFERRING - request, src_kv_list, _ = Sender._make_agent_request( - agent_args, is_aux=False, device_id=self._device_id + def _deliver_kv_to_agent(self, write_meta: WriteMeta): + assert len(write_meta.src_ptrs) == len(write_meta.dst_ptrs) == len(write_meta.sizes), ( + f"WriteMeta ptr/size mismatch for unique_rid={write_meta.unique_rid}" ) - skip_send = len(src_kv_list) == 0 - logger.debug(f"Submitting transfer request to transfer agent: {request}") - agent_handler = None - if task._perf_timer is not None: - task._perf_timer.record_transfer_start(agent_args.peer_rank) - if not skip_send: - agent_handler = self._agent.submit_transfer_requests(request) - - sync_status = "SUCCESS" - if not skip_send and not agent_handler.wait(): - sync_status = "FAILED" - agent_args.future_for_task.set_exception(RuntimeError("Transfer failed")) - task.status = TaskStatus.ERROR - session.state.status = SessionStatus.ERROR - if task._perf_timer is not None: - task._perf_timer.record_transfer_end(agent_args.peer_rank) + session = self._get_session(write_meta.unique_rid) + if session is None: + msg = ( + f"_deliver_kv_to_agent: TxSession {write_meta.unique_rid} not found or already GC'd" + ) + logger.error(msg) + if not write_meta.task_future.done(): + write_meta.task_future.set_exception(RuntimeError(msg)) + return + assert write_meta.slice_id is not None + task = session.kv_tasks[write_meta.slice_id] + timer = task._perf_timer + if timer: + timer.record_push_end(write_meta.peer_rank) + if session.status == SessionStatus.ERROR: + logger.warning( + f"_deliver_kv_to_agent: session {write_meta.unique_rid} already in ERROR state, skipping" + ) + return + task.status = TaskStatus.TRANSFERRING - messenger = self._get_or_connect_dealer(peer_endpoint) + agent_result = AgentResult.SUCCESS + if timer: + timer.record_transfer_start(write_meta.peer_rank) + if write_meta.src_ptrs: + request = Sender._make_agent_request(write_meta, device_id=self._device_id) + if not self._agent.submit_transfer_requests(request).wait(): + agent_result = AgentResult.FAILED + if not write_meta.task_future.done(): + write_meta.task_future.set_exception( + RuntimeError(f"KV transfer failed for request {write_meta.unique_rid}") + ) + task.status = TaskStatus.ERROR + if timer: + timer.record_transfer_end(write_meta.peer_rank) ## TODO: just last slice need to send task state? - messenger.send( + self._get_or_connect_dealer(write_meta.peer_endpoint).send( [ - MessageType.TASK_STATUS, + MessageType.KV_AGENT_RESULT, str(self._instance_rank).encode("ascii"), - str(unique_rid).encode("ascii"), - str(slice_id).encode("ascii"), - str(agent_args.is_last_slice).encode("ascii"), - sync_status.encode("ascii"), + str(write_meta.unique_rid).encode("ascii"), + str(write_meta.slice_id).encode("ascii"), + str(write_meta.is_last_slice).encode("ascii"), + agent_result.value.encode("ascii"), ] ) - curr = task.transferred_count - task.transferred_count = curr + 1 - if task._perf_timer is not None: - task._perf_timer.record_task_end(agent_args.peer_rank) - task.print_perf_info(agent_args.peer_rank) - if task.transferred_count > agent_args.expected_transfers: - agent_args.future_for_task.set_exception( - RuntimeError( - f"Session {unique_rid} has more than {agent_args.expected_transfers} transfers" - ) + task.transferred_count += 1 + if timer: + timer.record_task_end(write_meta.peer_rank) + ri = self._registrar.self_rank_info + task.print_perf_info(write_meta.peer_rank, ri.instance_name, ri.instance_rank) + if task.transferred_count > write_meta.expected_transfers: + session.set_exception( + f"KV slice {write_meta.slice_id} received more than {write_meta.expected_transfers} transfers" ) - # TODO: set exception for the session ? - session.state.status = SessionStatus.ERROR - elif task.transferred_count == agent_args.expected_transfers: - # TODO avoid set_result if tranfser failed since it has been set exception - agent_args.future_for_task.set_result(sync_status) - task.status = TaskStatus.TRANSFERRED - session.state.finished_tasks.append(slice_id) - if agent_args.is_last_slice: - session.state.status = SessionStatus.TRANSFERRED + elif task.transferred_count == write_meta.expected_transfers: + if write_meta.task_future.done(): + task.status = TaskStatus.ERROR + session.set_exception( + f"KV slice {write_meta.slice_id} future already resolved on completion" + ) + else: + write_meta.task_future.set_result(AgentResult.SUCCESS) + task.status = TaskStatus.TRANSFERRED logger.debug( - f"deliver_kv_to_agent completed: unique_rid={agent_args.unique_rid}, " - f"slice_id={slice_id}, sync_status={sync_status}" + f"deliver_kv_to_agent completed: unique_rid={write_meta.unique_rid}, " + f"slice_id={write_meta.slice_id}, agent_result={agent_result}" ) - @nvtx_range("submit_send_aux_to_agent") - def _deliver_aux_to_agent(self, agent_args: WriteMeta): - # TODO: submit the aux data task to the transfer agent - assert agent_args.is_only_aux is True - # assert agent_args.src_aux_ptrs is not None - - session = self._get_tx_session(agent_args.unique_rid) - assert session is not None, f"cannot get session for unique_rid {agent_args.unique_rid}" - if session._aux_task._perf_timer is not None: - session._aux_task._perf_timer.record_push_end(agent_args.peer_rank) - skip_send = len(agent_args.src_aux_ptrs) == 0 - sync_status = "SUCCESS" - agent_handler = None - if not skip_send: - request, _, _ = Sender._make_agent_request( - agent_args, is_aux=True, device_id=self._device_id - ) - agent_handler = self._agent.submit_transfer_requests(request) - - if session._aux_task._perf_timer is not None: - session._aux_task._perf_timer.record_transfer_start(agent_args.peer_rank) - if not agent_handler.wait(): - sync_status = "FAILED" - agent_args.future_for_task.set_exception(RuntimeError("Transfer failed")) - session.state.status = SessionStatus.ERROR - - if session._aux_task._perf_timer is not None: - session._aux_task._perf_timer.record_transfer_end(agent_args.peer_rank) - messenger = self._get_or_connect_dealer(agent_args.peer_endpoint) - messenger.send( + @nvtx_range("_deliver_aux_to_agent") + def _deliver_aux_to_agent(self, write_meta: WriteMeta): + session = self._get_session(write_meta.unique_rid) + if session is None: + msg = f"_deliver_aux_to_agent: TxSession {write_meta.unique_rid} not found or already GC'd" + logger.error(msg) + if not write_meta.task_future.done(): + write_meta.task_future.set_exception(RuntimeError(msg)) + return + aux_task = session.aux_task + assert aux_task is not None, f"aux_task is None for session {write_meta.unique_rid}" + timer = aux_task._perf_timer + if timer: + timer.record_push_end(write_meta.peer_rank) + + agent_result = AgentResult.SUCCESS + if write_meta.src_ptrs: + request = Sender._make_agent_request(write_meta, device_id=self._device_id) + if timer: + timer.record_transfer_start(write_meta.peer_rank) + if not self._agent.submit_transfer_requests(request).wait(): + agent_result = AgentResult.FAILED + session.set_exception("aux transfer agent request failed") + if timer: + timer.record_transfer_end(write_meta.peer_rank) + + self._get_or_connect_dealer(write_meta.peer_endpoint).send( [ - MessageType.AUX_SEND_STATUS, + MessageType.AUX_AGENT_RESULT, str(self._instance_rank).encode("ascii"), - str(agent_args.unique_rid).encode("ascii"), - sync_status.encode("ascii"), + str(write_meta.unique_rid).encode("ascii"), + agent_result.value.encode("ascii"), ] ) - aux_task = session._aux_task - aux_task._transferred_count += 1 - if aux_task._perf_timer is not None: - aux_task._perf_timer.record_task_end(agent_args.peer_rank) - aux_task.print_perf_info(agent_args.peer_rank) - if aux_task._transferred_count == agent_args.expected_transfers: - aux_task.future.set_result(sync_status) - aux_task.status = TaskStatus.AUX_TRANSFERRED - session.state.status = SessionStatus.AUX_TRANSFERRED - elif aux_task._transferred_count > agent_args.expected_transfers: - aux_task.future.set_exception( - RuntimeError( - f"Session {agent_args.unique_rid} has more than {agent_args.expected_transfers} transfers" - ) + + aux_task._transfer_count += 1 + if timer: + timer.record_task_end(write_meta.peer_rank) + ri = self._registrar.self_rank_info + aux_task.print_perf_info(write_meta.peer_rank, ri.instance_name, ri.instance_rank) + if aux_task._transfer_count == write_meta.expected_transfers: + if aux_task.future.done(): + aux_task.status = TaskStatus.ERROR + session.set_exception("aux future already resolved on completion") + else: + aux_task.future.set_result(AgentResult.SUCCESS) + aux_task.status = TaskStatus.TRANSFERRED + elif aux_task._transfer_count > write_meta.expected_transfers: + session.set_exception( + f"aux task received more than {write_meta.expected_transfers} transfers" ) - session.state.status = SessionStatus.ERROR - def dispatch_task(self, task: KVSendTask | AuxSendTask, req_info: Optional[RecvReqInfo] = None): - if not task.is_active(): - logger.debug(f"TxTask {task} is not active, skipping dispatch") - return + @staticmethod + def _filter_kv_blocks(src_block_ids, dst_block_ids) -> tuple[list[int], list[int]]: + # TODO: filter the kv block_ids according to the peer_overlap + return src_block_ids, dst_block_ids + + @nvtx_range("_build_kv_write_meta") + def _build_kv_write_meta(self, task: KVSendTask, req_info: RecvReqInfo) -> WriteMeta: + peer_ri = self._registrar.get_peer_rank_info(req_info.instance_name, req_info.instance_rank) + timer = task._perf_timer + if timer: + timer.record_prepare_args_start(peer_ri.instance_rank) + targets = self._registrar.get_peer_overlap(peer_ri, peer_ri.dp_rank) + expected_transfers = len(targets.ranks) + + src_frags: List[int] = [] + dst_frags: List[int] = [] + kv_sizes: List[int] = [] + dst_device_id = None + if self._registrar.should_send_kv(targets, peer_ri): + dst_device_id = peer_ri.device_id + extractor = self._registrar.self_extractor + 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] + + if len(src_block_ids) + 1 == len(dst_block_ids): + # FIXME: this is a temporary solution, need to be fixed for the draft tokens + logger.warning( + "src_block_num is one less than dst_block_num, maybe it is due to draft tokens," + " remove the last block from dst_block_ids " + ) + dst_block_ids = dst_block_ids[:-1] + src_block_ids, dst_block_ids = Sender._filter_kv_blocks( + src_block_ids, dst_block_ids + ) + + src_region = extractor.extract( + src_block_ids, layer_group_id=self_lg, pool_idx=self_pi + ) + dst_region = peer_extractor.extract( + dst_block_ids, layer_group_id=peer_lg, pool_idx=peer_pi + ) + mapper = self._registrar.get_kv_map(peer_ri, (self_lg, self_pi), (peer_lg, peer_pi)) + region_pair = mapper.map(src_region, dst_region) + region_pairs = region_pair if isinstance(region_pair, list) else [region_pair] + for rp in region_pairs: + src_frags.extend(rp.src.memory.ptrs) # type: ignore[attr-defined] + dst_frags.extend(rp.dst.memory.ptrs) # type: ignore[attr-defined] + frag_size = rp.src.memory.bytes_per_region # type: ignore[attr-defined] + kv_sizes.extend([frag_size] * len(rp.src.memory.ptrs)) # type: ignore[attr-defined] + + if timer: + timer.record_prepare_args_end(peer_ri.instance_rank) + timer.record_transfer_sizes(peer_ri.instance_rank, sum(kv_sizes), len(dst_frags)) + + return WriteMeta( + task_future=task.future, + src_ptrs=src_frags, + dst_ptrs=dst_frags, + sizes=kv_sizes, + dst_device_id=dst_device_id, + expected_transfers=expected_transfers, + peer_name=peer_ri.instance_name + str(peer_ri.instance_rank), + peer_rank=peer_ri.instance_rank, + peer_endpoint=peer_ri.self_endpoint, + unique_rid=task._unique_rid, + slice_id=task.slice_id, + is_last_slice=task._slice.is_last_slice, + ) + + def _build_aux_write_meta(self, task: AuxSendTask, req_info: RecvReqInfo) -> WriteMeta: + peer_ri = self._registrar.get_peer_rank_info(req_info.instance_name, req_info.instance_rank) + timer = task._perf_timer + if timer: + timer.record_prepare_args_start(peer_ri.instance_rank) + expected_transfers = len(self._registrar.get_peer_overlap(peer_ri, peer_ri.dp_rank).ranks) + + src_ptrs, dst_ptrs, sizes = [], [], [] + if self._registrar.should_send_aux(peer_ri): + src_aux_meta = self._registrar.self_rank_info.aux_meta + peer_aux_meta = peer_ri.aux_meta + assert src_aux_meta is not None + assert peer_aux_meta is not None + peer_slot = req_info.aux_slot + assert peer_slot is not None, f"aux_slot is None for request {req_info.unique_rid}" + assert task._slot is not None + src_ptrs = [ + ptr + item_size * task._slot + for ptr, item_size in zip(src_aux_meta.ptrs, src_aux_meta.item_sizes) + ] + dst_ptrs = [ + ptr + item_size * peer_slot + for ptr, item_size in zip(peer_aux_meta.ptrs, peer_aux_meta.item_sizes) + ] + sizes = list(src_aux_meta.item_sizes) + + if timer: + timer.record_prepare_args_end(peer_ri.instance_rank) + timer.record_transfer_sizes(peer_ri.instance_rank, sum(sizes), len(src_ptrs)) + + return WriteMeta( + task_future=task.future, + src_ptrs=src_ptrs, + dst_ptrs=dst_ptrs, + sizes=sizes, + expected_transfers=expected_transfers, + peer_name=req_info.instance_name + str(req_info.instance_rank), + peer_rank=req_info.instance_rank, + peer_endpoint=peer_ri.self_endpoint, + unique_rid=task._unique_rid, + meta_type=WriteMetaType.AUX, + ) - def dispatch_task_with_req_info(info: RecvReqInfo): + def dispatch_task( + self, + task: KVSendTask | AuxSendTask, + req_info_snapshot: Optional[dict] = None, + ): + # req_info_snapshot may be pre-fetched under session.lock by the caller to keep the + # critical section small. When not provided, we fetch it here (legacy / standalone path). + if req_info_snapshot is None: + req_info_snapshot = dict(self._get_req_info(task._unique_rid) or {}) + for info in req_info_snapshot.values(): if task._perf_timer is not None: task._perf_timer.record_task_start(info.instance_rank) - trans_meta = task._create_write_meta(info) + if isinstance(task, KVSendTask): + trans_meta = self._build_kv_write_meta(task, info) + else: + trans_meta = self._build_aux_write_meta(task, info) if task._perf_timer is not None: task._perf_timer.record_push_start(trans_meta.peer_rank) - self.submit_task(trans_meta) - - if req_info is None: - all_req_infos = self._peer_reqs.get_req_info(task._unique_rid) - for _, peer_req_info in all_req_infos.items(): - dispatch_task_with_req_info(peer_req_info) - else: - dispatch_task_with_req_info(req_info) + self._enqueue(trans_meta) def _start_listener(self): def handle_message(messages: list[bytes]): @@ -773,16 +608,23 @@ def handle_message(messages: list[bytes]): case MessageType.TERMINATION: return False case MessageType.REQUEST_DATA: - self._handle_request_data(send_id, msg) + try: + self._respond_with_kv(send_id, msg) + except Exception as e: + logger.error(f"Sender: error handling REQUEST_DATA: {e}") case MessageType.REGISTER_RANK_INFO: - self._register_peer_rank(send_id, msg) + try: + self._register_peer_rank(send_id, msg) + except Exception as e: + logger.error(f"Sender: error handling REGISTER_RANK_INFO: {e}") case _: - raise ValueError(f"Sender received unknown message type: {msg[0]}") + logger.error(f"Sender received unknown message type: {msg[0]}") - self._messenger.start_listener(handle_message, _handle_error) + self._messenger.start_listener(handle_message) - def _register_peer_rank(self, send_id: bytes, message: list[bytes]): + def _register_peer_rank(self, _send_id: bytes, message: list[bytes]): ri: RankInfo = RankInfo.from_bytes(message[1]) + self._registrar.register(ri.instance_name, ri.instance_rank, ri) agent_name = ri.instance_name + str(ri.instance_rank) @@ -796,36 +638,48 @@ def _register_peer_rank(self, send_id: bytes, message: list[bytes]): f"Completed handling REGISTER_RANK_INFO for instance='{ri.instance_name}', rank={ri.instance_rank}" ) - @nvtx_range("_handle_request_data") - def _handle_request_data(self, send_id: bytes, message: list[bytes]): - # each context send task may respond to multiple req_infos which is from different gen ranks. - # we support that tx_session send is called before or after getting the req_info. - # which means we need to support three cases: - # 1. the tx_session send is called before getting the all req_infos. - # 2. the tx_session send is called after getting the all req_infos. - # 3. the tx_session send is called after getting some req_infos and before getting the other req_infos. - # response_with_kv and function tx_session send both may trigger the submit task, - # use session.lock to serialize the submit task in both functions to avoid duplication. - # use self._sessions_lock to ensure the session will not be inserted between - # self._save_peer_req_info and get_tx_session - req_info: RecvReqInfo = RecvReqInfo.from_bytes(message[1]) + @nvtx_range("_respond_with_kv") + def _respond_with_kv(self, _send_id: bytes, message: list[bytes]): + # A session's KV send may race with incoming req_infos from multiple gen ranks. + # _sessions_lock guards against session insertion between _save_peer_req_info and + # _get_session; session.lock serializes _enqueue calls from both paths. + info: RecvReqInfo = RecvReqInfo.from_bytes(message[1]) with self._sessions_lock: - session: Optional[TxSession] = self._get_tx_session(req_info.unique_rid) + session = self._get_session(info.unique_rid) if session is None: - self._peer_reqs.add_req_info(req_info.unique_rid, req_info.instance_rank, req_info) + self._save_peer_req_info(info) return - with session._lock: - self._peer_reqs.add_req_info(req_info.unique_rid, req_info.instance_rank, req_info) - if self._has_all_peer_req_infos(req_info): - if session.state.status == SessionStatus.INIT: - session.state.status = SessionStatus.READY - # Dispatch incrementally as each peer's req_info arrives (case 3 above). - # delay dispatching for gen-first request until respond_and_send_async is called - if session.disagg_params.schedule_style != DisaggScheduleStyle.GENERATION_FIRST: - session.dispatch_all_tasks(req_info=req_info) + with session.lock: + self._save_peer_req_info(info) + tasks = list(session.kv_tasks) + for task in tasks: + if task._perf_timer is not None: + task._perf_timer.record_task_start(info.instance_rank) + trans_meta = self._build_kv_write_meta(task, info) + if task._perf_timer is not None: + task._perf_timer.record_push_start(trans_meta.peer_rank) + self._enqueue(trans_meta) + + def _get_or_connect_dealer(self, endpoint: Optional[str]): + if endpoint is None: + raise ValueError("Sender: peer endpoint is None; peer may not have registered yet") + if endpoint not in self._dealers: + self._dealers[endpoint] = ZMQMessenger(mode="DEALER", endpoint=endpoint) + return self._dealers[endpoint] + + def _save_peer_req_info(self, peer_transfer_req_info: RecvReqInfo): + req_info = peer_transfer_req_info + self._add_req_info(req_info.unique_rid, req_info.instance_rank, req_info) + 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) + if self._is_req_ready(req_info.unique_rid, expected_transfers): + if req_info.unique_rid in self._sessions: + session = self._get_session(req_info.unique_rid) + if session is not None and not session.receiver_ready: + session.receiver_ready = True def has_all_peer_req_infos(self, unique_rid: int) -> bool: - req_info = self._peer_reqs.get_first_req_info(unique_rid) + req_info = self._get_first_req_info(unique_rid) if req_info: return self._has_all_peer_req_infos(req_info) return False @@ -834,234 +688,222 @@ 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._peer_reqs.is_ready(req_info.unique_rid, expected_transfers) - - def _get_or_connect_dealer(self, endpoint: str): - if endpoint is None: - raise ValueError("endpoint is None") - if endpoint not in self._dealers: - self._dealers[endpoint] = ZMQMessenger(mode="DEALER", endpoint=endpoint) - return self._dealers[endpoint] + return self._is_req_ready(req_info.unique_rid, expected_transfers) def clear_session(self, unique_rid: int): - """Clear session-related resources from Sender. - - Args: - unique_rid: The unique request ID of the session to clear - """ with self._sessions_lock: - if unique_rid in self._tx_sessions: - del self._tx_sessions[unique_rid] - self._peer_reqs.remove_req_info(unique_rid) + if unique_rid in self._sessions: + del self._sessions[unique_rid] + self._remove_req_info(unique_rid) def shutdown(self): - if self._closed: + if self._shutdown: return - self._closed = True + self._shutdown = True # Stop all worker threads by sending None to each queue - if hasattr(self, "_send_task_queues"): - for q in self._send_task_queues: - q.put(None) - if hasattr(self, "_worker_threads"): - for t in self._worker_threads: - t.join(timeout=5) - # Invalidate all loaded remote agents to release fabric/POSIX FD resources - if hasattr(self, "_loaded_remote_agents"): - for agent_name in self._loaded_remote_agents: - try: - self._agent.invalidate_remote_agent(agent_name) - except Exception as e: - logger.warning( - f"Failed to invalidate remote agent '{agent_name}' during shutdown: {e}" - ) - self._loaded_remote_agents.clear() + for q in self._send_task_queues: + q.put(None) + for t in self._worker_threads: + t.join(timeout=5) + # Invalidate all loaded remote agents to release fabric/POSIX FD resources + for agent_name in self._loaded_remote_agents: + try: + self._agent.invalidate_remote_agent(agent_name) + except Exception as e: + logger.warning( + f"Failed to invalidate remote agent '{agent_name}' during shutdown: {e}" + ) + self._loaded_remote_agents.clear() + for dealer in self._dealers.values(): + try: + dealer.stop() + except Exception as e: + logger.warning(f"Failed to stop dealer during Sender shutdown: {e}") + self._dealers.clear() self._messenger.stop() def __del__(self): try: self.shutdown() except Exception as e: - logger.error(f"Exception in Sender.__del__: {e}") + logger.warning(f"Sender.__del__: exception during shutdown: {e}") def __enter__(self): return self - def __exit__(self, exc_type, exc_val, exc_tb): + def __exit__(self, _exc_type, _exc_val, _exc_tb): self.shutdown() class TxSession(TxSessionBase): def __init__( self, - request: LlmRequest, + request_id: int, + params: DisaggregatedParams, sender: Sender, - aux_slot: Optional[AuxSlot], + aux_slot: Optional[int], + aux_buffer: Optional[AuxBuffer] = None, ): - super().__init__(sender, request) - self._aux_slot = aux_slot + 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.receiver_ready: bool = False + self.kv_tasks = [] + self.aux_task = None + self.lock = threading.Lock() + + self._exception: Optional[Exception] = None self._closed = False - self._kv_tasks = [] - self._aux_task = None - self._lock = threading.Lock() - # setup_session registers this session, making it visible to the - # listener thread. All instance attributes that the listener may - # access (_kv_tasks, _lock, etc.) must be initialised BEFORE this call - # to avoid a race where the listener sees the session before its - # attributes are ready. + # Must be last: makes session visible to listener thread, + # so all attributes above must be initialized first. self._sender.setup_session(self) @property - def aux_slot(self) -> AuxSlot: - return self._aux_slot - - def send(self, slice: KVSlice) -> TaskIdType: - with self._lock: - slice_id = len(self._kv_tasks) - task = KVSendTask(slice, self.unique_rid, slice_id, self._sender._registrar) - self._kv_tasks.append(task) - self._sender.dispatch_task(task) - return task.slice_id + def status(self) -> SessionStatus: + if self._exception is not None or any(t.status == TaskStatus.ERROR for t in self.kv_tasks): + return SessionStatus.ERROR + kv_all_transferred = bool(self.kv_tasks) and all( + t.status == TaskStatus.TRANSFERRED for t in self.kv_tasks + ) + if kv_all_transferred: + if self.aux_task is not None and self.aux_task.status == TaskStatus.TRANSFERRED: + return SessionStatus.FULLY_TRANSFERRED + return SessionStatus.KV_TRANSFERRED + if self.kv_tasks and any(t.status == TaskStatus.TRANSFERRING for t in self.kv_tasks): + return SessionStatus.TRANSFERRING + return SessionStatus.READY if self.receiver_ready else SessionStatus.INIT + + def send(self, slice: KVSlice) -> concurrent.futures.Future: + with self.lock: + params = self._base_args.params + slice_id = len(self.kv_tasks) + task = KVSendTask(slice, params, slice_id) + self.kv_tasks.append(task) + req_info_snapshot = dict(self._sender._get_req_info(task._unique_rid) or {}) + self._sender.dispatch_task(task, req_info_snapshot) + return task.future def send_aux(self) -> AuxSendTask: - self.pack_aux() - with self._lock: - slot = self._aux_slot.id - task = AuxSendTask(self.unique_rid, slot, self._sender._registrar) - self._aux_task = task - self._sender.dispatch_task(task) - return task - - def pack_aux(self) -> None: - self._aux_slot.buffer.fill_slot(self._aux_slot.id, self.request) - - def poll_task(self, id: TaskIdType) -> SessionStatus: - return self._kv_tasks[id].state - - def dispatch_all_tasks(self, req_info: Optional[RecvReqInfo] = None): - # call with lock held - for task in self._kv_tasks: - if task.is_active(): - self._sender.dispatch_task(task, req_info) - if self._aux_task: - self._sender.dispatch_task(self._aux_task, req_info) - - def wait_complete( - self, task_id: TaskIdType, wait_aux: bool = True, timeout_ms: int = -1 - ) -> bool: - timeout_s = timeout_ms / 1000.0 if timeout_ms > 0 else None - kv_result = self._kv_tasks[task_id].future.result(timeout=timeout_s) == "SUCCESS" - if wait_aux and self._aux_task: - aux_result = self._aux_task.future.result(timeout=timeout_s) == "SUCCESS" - return kv_result and aux_result - else: - return kv_result + with self.lock: + params = self._base_args.params + task = AuxSendTask(params, self.aux_slot) + self.aux_task = task + req_info_snapshot = dict(self._sender._get_req_info(task._unique_rid) or {}) + self._sender.dispatch_task(task, req_info_snapshot) + return task + + def pack_aux(self, request: LlmRequest) -> None: + """Fill the aux buffer slot with token data from the given request.""" + assert self._aux_buffer is not None, "No aux_buffer set for this session" + 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: + """Non-blocking check: has the transfer completed successfully?""" + status = self.status + if need_aux: + return status == SessionStatus.FULLY_TRANSFERRED + return status in (SessionStatus.KV_TRANSFERRED, SessionStatus.FULLY_TRANSFERRED) + + 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: + """Block until KV (and optionally aux) transfer finishes. + + 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 + if need_aux and self.aux_task is not None: + aux_status = self.aux_task.future.result(timeout=timeout) + if aux_status != AgentResult.SUCCESS: + return WaitResult.FAILED + return WaitResult.COMPLETED + except TimeoutError: + return WaitResult.TIMEOUT + except Exception: + return WaitResult.FAILED + + def set_exception(self, reason: str = ""): + msg = f"TxSession {self.disagg_request_id} exception" + if reason: + msg += f": {reason}" + self._exception = RuntimeError(msg) + for task in self.kv_tasks: + if not task.future.done(): + task.future.set_exception(self._exception) + if self.aux_task is not None and not self.aux_task.future.done(): + self.aux_task.future.set_exception(self._exception) + + @property + def exception(self) -> Optional[Exception]: + return self._exception def close(self): if getattr(self, "_closed", False): return self._closed = True - # Clear session from Sender's lookup dict. - # Do NOT null out _kv_tasks/_aux_task/_sender here: - # worker threads may still hold a strong reference to this session - # and need access to these fields to finish in-flight transfers. - # Resources are freed when the session object is garbage collected. + # 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: - self._sender.clear_session(self.unique_rid) + self._sender.clear_session(self.disagg_request_id) def __enter__(self): return self - def __exit__(self, exc_type, exc, tb): + def __exit__(self, _exc_type, _exc, _tb): self.close() def __del__(self): try: self.close() - except Exception: - pass + except Exception as e: + logger.warning(f"TxSession.__del__: exception during close: {e}") class KVRecvTask: def __init__( self, - unique_rid: int, - disagg_params: DisaggregatedParams, + unique_rid: Optional[int], kv_slice: KVSlice, slice_id: int, - peer_registrar: PeerRegistrar, - aux_slot: Optional[AuxSlot], + params: DisaggregatedParams, + aux_slot: Optional[int], ): - self.unique_rid = unique_rid - self.disagg_params = disagg_params + self.future = concurrent.futures.Future() + self.slice_id = slice_id + self.status = TaskStatus.INIT + self.expected_transfers = 0 + self.last_slice_count = 0 + + self._unique_rid = unique_rid self._kv_slice = kv_slice - self._slice_id = slice_id - self._registrar = peer_registrar - self._status = TaskStatus.INIT + self._params = params self._exception = None - self._future = concurrent.futures.Future() - self._first_transfer = False - self._expected_transfers = 0 - self._aux_slot_id = aux_slot.id if aux_slot else None + self._aux_slot = aux_slot self._perf_timer = PerfTimer() if perf_log_manager.enabled else None - @property - def status(self) -> TaskStatus: - return self._status - - @status.setter - def status(self, s: TaskStatus): - self._status = s - - @property - def future(self) -> concurrent.futures.Future: - return self._future - - @property - def slice_id(self) -> int: - return self._slice_id - - @property - def expected_transfers(self) -> int: - return self._expected_transfers - - @expected_transfers.setter - def expected_transfers(self, v: int): - self._expected_transfers = v - - def create_req_info(self) -> RecvReqInfo: - return RecvReqInfo( - sender_req_id=self.disagg_params.ctx_request_id, - instance_name=self._registrar.self_rank_info.instance_name, - instance_rank=self._registrar.self_rank_info.instance_rank, - block_ids_per_layer_groups=self._kv_slice.block_ids_per_layer_groups, - unique_rid=self.unique_rid, - aux_slot=self._aux_slot_id, - ) - - def make_read_meta(self, peer_ii, peer_dp_rank) -> ReadMeta: - peer_overlap = self._registrar.get_peer_overlap(peer_ii, peer_dp_rank) - if not self._first_transfer: - self._first_transfer = True - self._expected_transfers = len(peer_overlap.ranks) - return ReadMeta( - slice_id=self._slice_id, - unique_rid=self.unique_rid, - target_ranks=peer_overlap.ranks, - ) - - def print_perf_info(self, peer_rank: int): - ri = self._registrar.self_rank_info + def print_perf_info(self, peer_rank: int, instance_name: str, instance_rank: int): + if self._perf_timer is None: + return + assert self._unique_rid is not None perf_log_manager.log_recv_task_perf( - self.unique_rid, + self._unique_rid, peer_rank, - ri.instance_name, - ri.instance_rank, + instance_name, + instance_rank, self._perf_timer, ) -class Receiver: +class Receiver(ReceiverBase): def __init__( self, peer_registrar: PeerRegistrar, @@ -1075,35 +917,40 @@ def __init__( self._sender_ep_instance_map = {} self._messenger = ZMQMessenger(mode="ROUTER") + self._sessions = {} # unique_rid -> RxSession + self._sessions_lock = threading.Lock() + self._shutdown = False + self._start_listener() logger.info(f"Receiver init with endpoint: {self._messenger.endpoint}") - self._rx_sessions = {} # unique_rid -> RxSession - self._closed = False - @property def endpoint(self): return self._messenger.endpoint def shutdown(self): - if getattr(self, "_closed", False): + if getattr(self, "_shutdown", False): return - self._closed = True + self._shutdown = True + for dealer in self._dealers.values(): + try: + dealer.stop() + except Exception as e: + logger.warning(f"Failed to stop dealer during Receiver shutdown: {e}") + self._dealers.clear() self._messenger.stop() def clear_session(self, unique_rid: int): - """Clear session-related resources from Receiver. - - Args: - unique_rid: The unique request ID of the session to clear - """ - self._rx_sessions.pop(unique_rid, None) + with self._sessions_lock: + self._sessions.pop(unique_rid, None) def setup_session(self, rx_session: RxSessionBase): - self._rx_sessions[rx_session.unique_rid] = weakref.ref(rx_session) + with self._sessions_lock: + self._sessions[rx_session.disagg_request_id] = weakref.ref(rx_session) - def _get_rx_session(self, unique_rid: int) -> RxSessionBase: - session_ref = self._rx_sessions.get(unique_rid) + def _get_session(self, unique_rid: Optional[int]) -> Optional["RxSession"]: + with self._sessions_lock: + session_ref = self._sessions.get(unique_rid) if session_ref is None: return None session = session_ref() @@ -1112,53 +959,87 @@ def _get_rx_session(self, unique_rid: int) -> RxSessionBase: return None return session + def _build_recv_req_info(self, task: KVRecvTask) -> RecvReqInfo: + self_ri = self._registrar.self_rank_info + assert task._params.ctx_request_id is not None, ( + f"ctx_request_id is None for task unique_rid={task._unique_rid}" + ) + assert task._unique_rid is not None, "KVRecvTask unique_rid is None" + return RecvReqInfo( + sender_req_id=task._params.ctx_request_id, + instance_name=self_ri.instance_name, + instance_rank=self_ri.instance_rank, + block_ids_per_layer_groups=task._kv_slice.block_ids_per_layer_groups, + unique_rid=task._unique_rid, + aux_slot=task._aux_slot, + ) + def dispatch_task(self, task: KVRecvTask): - disagg_params = task.disagg_params - receiver_req = task.create_req_info() - sender_dp_rank = disagg_params.ctx_dp_rank + params = task._params + logger.debug(f"Preparing async data transfer request for disagg_params={params}") + receiver_req = self._build_recv_req_info(task) + sender_dp_rank = params.ctx_dp_rank if sender_dp_rank is None: - raise ValueError("sender_dp_rank is None") - peer_infos: RankInfo = self._get_sender_info(disagg_params) - agent_args = task.make_read_meta(peer_infos, sender_dp_rank) - session = self._get_rx_session(agent_args.unique_rid) - session._kv_tasks[agent_args.slice_id].status = TaskStatus.TRANSFERRING - for rank in agent_args.target_ranks: + raise ValueError( + f"ctx_dp_rank is None for request {task._unique_rid}; " + "disaggregated params may be missing context rank info" + ) + peer_infos: RankInfo = self._get_sender_info(params) + peer_overlap = self._registrar.get_peer_overlap(peer_infos, sender_dp_rank) + task.expected_transfers = len(peer_overlap.ranks) + session = self._get_session(task._unique_rid) + if session is None: + raise RuntimeError( + f"dispatch_task: RxSession {task._unique_rid} not found; " + "session may have been closed before dispatch" + ) + session.mark_transferring(task.slice_id) + for rank in peer_overlap.ranks: if task._perf_timer is not None: task._perf_timer.record_task_start(rank) self._request_sender_data(peer_infos.sender_endpoints[rank], receiver_req) return - def _need_register_peer_in_first_request(self, params: DisaggregatedParams) -> bool: - return params.ctx_info_endpoint not in self._sender_ep_instance_map + @staticmethod + def _extract_info_endpoint(params: DisaggregatedParams) -> Optional[str]: + ep = params.ctx_info_endpoint + if isinstance(ep, list): + return ep[0] if ep else None + return ep # str (backward compat) + + def _should_register_peer(self, params: DisaggregatedParams) -> bool: + endpoint = self._extract_info_endpoint(params) + return endpoint not in self._sender_ep_instance_map def _get_or_connect_dealer(self, endpoint: Optional[str]): if endpoint is None: - raise ValueError("endpoint is None") + raise ValueError("Receiver: peer endpoint is None; peer may not have registered yet") if endpoint not in self._dealers: self._dealers[endpoint] = ZMQMessenger(mode="DEALER", endpoint=endpoint) return self._dealers[endpoint] def _get_sender_info(self, params: DisaggregatedParams) -> RankInfo: - if self._need_register_peer_in_first_request(params): - logger.info( - f"Registering peer in first request to endpoint '{params.ctx_info_endpoint}'" - ) - messenger = ZMQMessenger(mode="DEALER", endpoint=params.ctx_info_endpoint) - messenger.send([MessageType.REQUEST_INSTANCE_INFO]) - message = messenger.receive() - sender_info = RankInfo.from_bytes(message[0]) - messenger.stop() + info_endpoint = self._extract_info_endpoint(params) + if self._should_register_peer(params): + logger.info(f"Registering peer in first request to endpoint '{info_endpoint}'") + messenger = ZMQMessenger(mode="DEALER", endpoint=info_endpoint) + try: + messenger.send([MessageType.REQUEST_INSTANCE_INFO]) + message = messenger.receive() + sender_info = RankInfo.from_bytes(message[0]) + finally: + messenger.stop() for endpoint in sender_info.sender_endpoints: dealer = self._get_or_connect_dealer(endpoint) rank_info = self._registrar.self_rank_info dealer.send([MessageType.REGISTER_RANK_INFO, rank_info.to_bytes()]) - self._sender_ep_instance_map[params.ctx_info_endpoint] = sender_info + self._sender_ep_instance_map[info_endpoint] = sender_info return sender_info else: - return self._sender_ep_instance_map[params.ctx_info_endpoint] + return self._sender_ep_instance_map[info_endpoint] def _start_listener(self): def handle_message(messages: list[bytes]) -> bool: @@ -1167,235 +1048,273 @@ def handle_message(messages: list[bytes]) -> bool: match msg[0]: case MessageType.TERMINATION: return False - case MessageType.TASK_STATUS: - self._process_kv_task_status(send_id, msg) - case MessageType.AUX_SEND_STATUS: - self._process_aux_state(send_id, msg) + case MessageType.KV_AGENT_RESULT: + try: + self._process_kv_agent_result(send_id, msg) + except Exception as e: + logger.error(f"Receiver: error handling KV_AGENT_RESULT: {e}") + case MessageType.AUX_AGENT_RESULT: + try: + self._process_aux_agent_result(send_id, msg) + except Exception as e: + logger.error(f"Receiver: error handling AUX_AGENT_RESULT: {e}") case _: - raise ValueError(f"Receiver received unknown message type: {msg[0]}") + logger.error(f"Receiver received unknown message type: {msg[0]}") + return True - self._messenger.start_listener(handle_message, _handle_error) + self._messenger.start_listener(handle_message) - def _process_kv_task_status(self, send_id: bytes, message: list[bytes]): - msg_type, peer_rank, unique_rid, _, is_last_slice_str, status = decode_message(message) + def _process_kv_agent_result(self, _send_id: bytes, message: list[bytes]): + msg_type, peer_rank, unique_rid, slice_id_str, is_last_slice_str, status = decode_message( + message + ) peer_rank = int(peer_rank) unique_rid = int(unique_rid) - assert msg_type.encode("ascii") == MessageType.TASK_STATUS - session = self._get_rx_session(unique_rid) - if session is not None: - session.process_kv_task_status(peer_rank, is_last_slice_str == "True", status) - else: - logger.warning(f"RxSession {unique_rid} not found when processing kv task status") + slice_id = int(slice_id_str) + if msg_type.encode("ascii") != MessageType.KV_AGENT_RESULT: + logger.error( + f"_process_kv_agent_result: unexpected msg_type={msg_type!r}, expected TASK_STATUS" + ) + return + session = self._get_session(unique_rid) + if session is None: + logger.warning( + f"_process_kv_agent_result: session {unique_rid} not found (already closed?), dropping status" + ) + return + session.process_kv_agent_result( + peer_rank, slice_id, is_last_slice_str == "True", AgentResult(status) + ) - def _process_aux_state(self, send_id: bytes, message: list[bytes]): - msg_type, peer_rank, unique_rid, status = decode_message(message) + def _process_aux_agent_result(self, _send_id: bytes, message: list[bytes]): + _msg_type, peer_rank, unique_rid, status = decode_message(message) peer_rank = int(peer_rank) unique_rid = int(unique_rid) - session = self._get_rx_session(unique_rid) - if session is not None: - session.process_aux_state(peer_rank, status) - else: - logger.warning(f"RxSession {unique_rid} not found when processing aux state") + session = self._get_session(unique_rid) + if session is None: + logger.warning( + f"_process_aux_agent_result: session {unique_rid} not found (already closed?), dropping status" + ) + return + session.process_aux_agent_result(peer_rank, AgentResult(status)) def _request_sender_data(self, endpoint: str, receiver_info: RecvReqInfo): + logger.debug( + f"Sending data request to endpoint '{endpoint}' with request info: {receiver_info}" + ) messenger = self._get_or_connect_dealer(endpoint) messenger.send([MessageType.REQUEST_DATA, receiver_info.to_bytes()]) def __del__(self): try: self.shutdown() - except Exception: - pass + except Exception as e: + logger.warning(f"Receiver.__del__: exception during shutdown: {e}") def __enter__(self): return self - def __exit__(self, exc_type, exc_val, exc_tb): + def __exit__(self, _exc_type, _exc_val, _exc_tb): self.shutdown() class RxSession(RxSessionBase): def __init__( self, - request: LlmRequest, + request_id: int, + params: DisaggregatedParams, receiver: Receiver, - aux_slot: Optional[AuxSlot], + aux_slot: Optional[int], + aux_buffer: Optional[AuxBuffer] = None, ): - super().__init__(receiver, request) - self._aux_slot = aux_slot - self._receiver.setup_session(self) + 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._exception: Optional[Exception] = None self._closed = False - self._kv_tasks = [] - self._last_slice_counts = 0 - self._aux_counts = 0 - # aux_slot can be 0; treat None as the only "no aux" marker. - self._aux_future = concurrent.futures.Future() if aux_slot is not None else None + self._kv_tasks: list[KVRecvTask] = [] + self._aux_count = 0 + self._aux_status: TaskStatus = TaskStatus.INIT + self._receiver.setup_session(self) @property - def aux_slot(self) -> AuxSlot: - return self._aux_slot - - def receive(self, slice: KVSlice) -> TaskIdType: + def status(self) -> SessionStatus: + if self._exception is not None or ( + self._kv_tasks and self._kv_tasks[0].status == TaskStatus.ERROR + ): + 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: + return SessionStatus.FULLY_TRANSFERRED + if task_status == TaskStatus.TRANSFERRED: + return SessionStatus.KV_TRANSFERRED + if task_status == TaskStatus.TRANSFERRING: + return SessionStatus.TRANSFERRING + return SessionStatus.INIT + + def mark_transferring(self, slice_id: int): + self._kv_tasks[slice_id].status = TaskStatus.TRANSFERRING + + def receive(self, slice: KVSlice) -> concurrent.futures.Future: + params = self._base_args.params slice_id = len(self._kv_tasks) task = KVRecvTask( - unique_rid=self.unique_rid, - disagg_params=self.disagg_params, - kv_slice=slice, - slice_id=slice_id, - peer_registrar=self._receiver._registrar, - aux_slot=self._aux_slot, + params.disagg_request_id, + slice, + slice_id, + params, + aux_slot=self.aux_slot, ) self._kv_tasks.append(task) self._receiver.dispatch_task(task) - return task.slice_id + return task.future - def process_kv_task_status(self, peer_rank: int, is_last_slice: bool, status: str): - task = self._kv_tasks[0] # receive task slice only support slice 0 - - if status == "SUCCESS": + def process_kv_agent_result( + self, peer_rank: int, slice_id: int, is_last_slice: bool, status: AgentResult + ): + task = self._kv_tasks[slice_id] + if status == AgentResult.SUCCESS: if is_last_slice: - self._last_slice_counts += 1 - if self._last_slice_counts == task.expected_transfers: - task.future.set_result("SUCCESS") + task.last_slice_count += 1 + if task.last_slice_count == task.expected_transfers: + if not task.future.done(): + task.future.set_result(AgentResult.SUCCESS) task.status = TaskStatus.TRANSFERRED - self.state.status = SessionStatus.TRANSFERRED - self.state.finished_tasks.append(0) + logger.debug( + f"KV transfer complete for request {self.request_id} slice {slice_id}" + ) if task._perf_timer is not None: task._perf_timer.record_task_end(peer_rank) - task.print_perf_info(peer_rank) - elif self._last_slice_counts > task.expected_transfers: - logger.error( - f"Session {self.unique_rid} has more than {task.expected_transfers} transfers" + ri = self._receiver._registrar.self_rank_info + task.print_perf_info(peer_rank, ri.instance_name, ri.instance_rank) + elif status == AgentResult.FAILED: + if not task.future.done(): + task.future.set_exception( + RuntimeError( + f"KV transfer failed for request {self.request_id} slice {slice_id}" ) - - elif status == "FAILED": - task.future.set_exception(RuntimeError(f"Task state: {status}")) + ) task.status = TaskStatus.ERROR - self.state.status = SessionStatus.ERROR else: - raise ValueError(f"Session received unknown task status: {status}") - - def process_aux_state(self, peer_rank: int, status: str): - task = self._kv_tasks[0] # receive task slice only support slice 0 - if status == "SUCCESS": - self._aux_counts += 1 - - if self._aux_counts == task.expected_transfers: - task.status = TaskStatus.AUX_TRANSFERRED - self.state.status = SessionStatus.AUX_TRANSFERRED - self.unpack_aux() - if self._aux_future is not None: - self._aux_future.set_result("SUCCESS") - elif self._aux_counts > task.expected_transfers: - logger.error( - f"Session {self.unique_rid} has more than {task.expected_transfers} transfers" + raise ValueError( + f"Session {self.request_id} received unknown task status: {status.value}" + ) + + 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 + if status == AgentResult.SUCCESS: + self._aux_count += 1 + + if self._aux_count == task.expected_transfers: + self._aux_status = TaskStatus.TRANSFERRED + elif self._aux_count > task.expected_transfers: + self._aux_status = TaskStatus.ERROR + self._exception = RuntimeError( + f"Session {self.request_id} received too many aux transfers" ) - self.state.status = SessionStatus.ERROR - if self._aux_future is not None: - self._aux_future.set_exception( - RuntimeError( - f"Task unexpected count: {self._aux_counts} > {task.expected_transfers}" - ) - ) - elif status == "FAILED": - self.state.status = SessionStatus.ERROR - if self._aux_future is not None: - self._aux_future.set_exception(RuntimeError(f"Task state: {status}")) + logger.error(str(self._exception)) + elif status == AgentResult.FAILED: + self._aux_status = TaskStatus.ERROR + self._exception = RuntimeError(f"Session {self.request_id} aux transfer failed") else: - if self._aux_future is not None: - self._aux_future.set_exception(RuntimeError(f"Task state: {status}")) raise ValueError( - f"Session {self.unique_rid} received unknown aux send status: {status}" + f"Session {self.request_id} received unknown aux send status: {status}" ) - def poll_task(self, id: TaskIdType) -> SessionStatus: - return self._kv_tasks[id].state - - def wait_complete( - self, kv_task_id: TaskIdType, wait_aux: bool = True, timeout_ms: int = -1 - ) -> bool: + @property + def exception(self) -> Optional[Exception]: + return self._exception + + def unpack_aux(self, request: LlmRequest) -> None: + """Read token data from the aux buffer slot into the given request.""" + assert self._aux_buffer is not None, "No aux_buffer set for this session" + assert self.aux_slot is not None, "No aux_slot set for this session" + first_gen_tokens, draft_tokens = self._aux_buffer.get_slot_tokens(self.aux_slot) + 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: + """Non-blocking check: has the transfer completed successfully?""" + status = self.status + if need_aux: + return status == SessionStatus.FULLY_TRANSFERRED + return status in (SessionStatus.KV_TRANSFERRED, SessionStatus.FULLY_TRANSFERRED) + + 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. + + 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). + Returns WaitResult.COMPLETED on full success, WaitResult.FAILED on error. + """ try: - timeout_s = timeout_ms / 1000.0 if timeout_ms > 0 else None - kv_result = self._kv_tasks[kv_task_id].future.result(timeout=timeout_s) == "SUCCESS" - if wait_aux and self._aux_future is not None: - aux_result = self._aux_future.result(timeout=timeout_s) == "SUCCESS" - return kv_result and aux_result - else: - return kv_result - except concurrent.futures.TimeoutError: - logger.warning( - f"RxSession {self.unique_rid} timed out waiting for completion " - f"after {timeout_ms} milliseconds." - ) - return False - except Exception as e: - logger.error(f"Exception in RxSession.wait_complete: {e}") - return False + kv_status = self._kv_tasks[0].future.result() + if kv_status != AgentResult.SUCCESS: + return WaitResult.FAILED + if 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: + return None # KV done, aux still in flight; re-poll next cycle + time.sleep(0.001) + return WaitResult.COMPLETED + except Exception: + return WaitResult.FAILED def close(self): if getattr(self, "_closed", False): return self._closed = True - # Clear session from Receiver's lookup dict. - # Do NOT null out _kv_tasks/_receiver or reset counters here: - # the listener thread may still hold a strong reference to this session - # and need access to these fields to finish processing in-flight status messages. - # Resources are freed when the session object is garbage collected. + # 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: - self._receiver.clear_session(self.unique_rid) - - def unpack_aux(self) -> None: - assert self.request is not None, "request is not set" - request = self.request - first_gen_tokens, draft_tokens = self._aux_slot.buffer.get_slot_tokens(self._aux_slot.id) - request.py_draft_tokens = draft_tokens - if request.context_phase_params is None: - request.context_phase_params = ContextPhaseParams( - first_gen_tokens=first_gen_tokens, - req_id=request.py_request_id, - opaque_state=b"", - draft_tokens=draft_tokens, - ctx_dp_rank=0, - disagg_info_endpoint="", - ) - else: - request.context_phase_params.first_gen_tokens = first_gen_tokens - request.context_phase_params.draft_tokens = draft_tokens - return request + self._receiver.clear_session(self.disagg_request_id) def __enter__(self): return self - def __exit__(self, exc_type, exc, tb): + def __exit__(self, _exc_type, _exc, _tb): self.close() def __del__(self): try: self.close() - except Exception: - pass + except Exception as e: + logger.warning(f"RxSession.__del__: exception during close: {e}") class RankInfoServer: def __init__(self, rank_info: RankInfo, addr: Optional[str] = None, port: Optional[int] = None): self._rank_info = rank_info + self._shutdown = False # must be set before _start_listener() so __del__ is safe if addr is None and port is None: endpoint = f"tcp://{get_local_ip()}:*" else: endpoint = f"tcp://{addr}:{port}" self._messenger = ZMQMessenger(mode="ROUTER", endpoint=endpoint) self._start_listener() - self._closed = False @property def endpoint(self) -> str: return self._messenger.endpoint def shutdown(self): - if self._closed: + if self._shutdown: return - self._closed = True + self._shutdown = True logger.debug("RankInfoServer.shutdown() called") self._messenger.stop() @@ -1407,27 +1326,29 @@ def handle_message(messages: list[bytes]) -> bool: case MessageType.TERMINATION: return False case MessageType.REQUEST_INSTANCE_INFO: - self._process_request_rank_info(send_id, msg) + try: + self._handle_rank_info_request(send_id, msg) + except Exception as e: + logger.error(f"RankInfoServer: error handling REQUEST_INSTANCE_INFO: {e}") case _: - raise ValueError( - f"Instance info server received unknown message type: {msg[0]}" - ) + logger.error(f"Instance info server received unknown message type: {msg[0]}") + return True self._messenger.start_listener(handle_message) - def _process_request_rank_info(self, send_id: bytes, _message: list[bytes]): + def _handle_rank_info_request(self, send_id: bytes, _message: list[bytes]): self._messenger.send([send_id, self._rank_info.to_bytes()]) def __del__(self): try: self.shutdown() - except Exception: - pass + except Exception as e: + logger.warning(f"RankInfoServer.__del__: exception during shutdown: {e}") def __enter__(self): return self - def __exit__(self, exc_type, exc_val, exc_tb): + def __exit__(self, _exc_type, _exc_val, _exc_tb): self.shutdown() @@ -1438,7 +1359,7 @@ def _deregister_registered_memory(transfer_agent, registered_memorys): while registered_memorys: register_memory = registered_memorys[0] try: - logger.info(f" transfer worker deregister memory {register_memory} ") + logger.info(f"Deregistering transfer memory: {register_memory}") transfer_agent.deregister_memory(register_memory) except Exception: logger.error("deregister memory failed in finalizer") @@ -1458,13 +1379,14 @@ def __init__( ): self._mapping = mapping - self._rank_info: RankInfo = None + 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 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) @@ -1473,9 +1395,8 @@ def __init__( self._kv_extractor = KVRegionExtractorV1(self._kv_cache_manager) self._peer_registrar = PeerRegistrar(self._rank_info, self._kv_extractor) - # NixlTransferAgent configuration from environment variables - # num_threads: number of dedicated threads for large batch transfers - # split_batch_size: batch size threshold to use dedicated threads (default 1024) + # 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: @@ -1495,7 +1416,6 @@ def __init__( 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.endpoint = self._sender.endpoint self._rank_info.self_endpoint = self._receiver.endpoint reg_snapshot = list(self._registered_mem) if self._registered_mem is not None else [] @@ -1504,95 +1424,95 @@ def __init__( ) def populate_instance_and_rank_info(self, endpoints: list[str], layer_num_per_pp: list[int]): + assert self._rank_info is not None self._rank_info.sender_endpoints = endpoints self._rank_info.layer_num_per_pp = 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=request, + request_id=request.py_request_id, + params=params, sender=self._sender, - aux_slot=self._aux_buffer.alloc_slot() if self._aux_buffer else None, + 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=request, + request_id=request.py_request_id, + params=params, receiver=self._receiver, - aux_slot=self._aux_buffer.alloc_slot() if self._aux_buffer else None, + aux_slot=aux_slot, + aux_buffer=self._aux_buffer, ) def clear_session(self, session: TxSession | RxSession): - if self._aux_buffer is not None and session.aux_slot is not None: - self._aux_buffer.free_slot(session.aux_slot.id) + 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): - rank = self._mapping.rank - - tp_size = self._mapping.tp_size - pp_size = self._mapping.pp_size - dp_size = self._mapping.dp_size - cp_size = self._mapping.cp_size - tp_rank = self._mapping.tp_rank - pp_rank = self._mapping.pp_rank - enable_attention_dp = self._mapping.enable_attention_dp - dp_rank = 0 - if enable_attention_dp: - dp_size = self._mapping.tp_size - dp_rank = tp_rank - cp_rank = self._mapping.cp_rank - is_mla = self._kv_cache_manager.kv_factor == 1 - self._kv_cache_manager.kv_factor - heads_num_per_rank = self._kv_cache_manager.num_kv_heads_per_layer[0] - tokens_per_block = self._kv_cache_manager.tokens_per_block - dims_per_head = self._kv_cache_manager.head_dim - element_bytes = get_size_in_bytes(1, self._kv_cache_manager.dtype) - layer_num_per_pp = [len(self._kv_cache_manager.pp_layers)] - sender_endpoints = [] - # Build page table from manager (supports V1 and V2) - page_table = build_page_table_from_manager(self._kv_cache_manager) + 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=rank, - tp_size=tp_size, - tp_rank=tp_rank, - pp_size=pp_size, - pp_rank=pp_rank, - dp_size=dp_size, - dp_rank=dp_rank, - cp_size=cp_size, - cp_rank=cp_rank, + 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=layer_num_per_pp, - sender_endpoints=sender_endpoints, + layer_num_per_pp=[len(kvm.pp_layers)], + sender_endpoints=[], server_endpoint="", self_endpoint="", transfer_engine_info=bytes(), attention=AttentionInfo( - kv_heads_per_rank=heads_num_per_rank, - tokens_per_block=tokens_per_block, - dims_per_head=dims_per_head, - element_bytes=element_bytes, + 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=is_mla, + is_mla=kvm.kv_factor == 1, ), aux_meta=self._aux_buffer.meta if self._aux_buffer is not None else None, - page_table=page_table, + # Build page table from manager (supports V1 and V2) + page_table=build_page_table_from_manager(kvm), ) 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) @@ -1619,10 +1539,11 @@ def _register_kv_cache(self): if memory_descs: reg_memory_desc = RegMemoryDescs("VRAM", memory_descs) self._agent.register_memory(reg_memory_desc) - logger.debug("Registered KV cache memory with transfer agent: %s", memory_descs) + logger.debug(f"Registered KV cache memory with transfer agent: {memory_descs}") self._registered_mem.append(reg_memory_desc) def _register_aux_buffer(self): + assert self._aux_buffer is not None aux_meta = self._aux_buffer.meta ptr_num = len(aux_meta.ptrs) ptr_descs = [] @@ -1633,27 +1554,51 @@ def _register_aux_buffer(self): logger.debug(f"Registered auxiliary buffer memory with transfer agent: {reg_memory_desc}") self._registered_mem.append(reg_memory_desc) + @property + def rank_info_server_endpoint(self) -> Optional[str]: + return self._rank_info_server.endpoint if self._rank_info_server is not None else None + + @property + def sender_endpoint(self) -> str: + return self._sender.endpoint + + @property + def page_table(self): + assert self._rank_info is not None + return self._rank_info.page_table + def shutdown(self): - if self._rank_info_server is not None: - self._rank_info_server.shutdown() - if self._sender is not None: - self._sender.shutdown() - if self._receiver is not None: - self._receiver.shutdown() + if getattr(self, "_shutdown", False): + return + self._shutdown = True + # Use getattr guards: __init__ may have failed partway, leaving some + # attributes unset. Without them, __del__ -> shutdown() raises + # AttributeError and ZMQ resources from already-created sub-objects + # are never cleaned up. + rank_info_server = getattr(self, "_rank_info_server", None) + if rank_info_server is not None: + rank_info_server.shutdown() + sender = getattr(self, "_sender", None) + if sender is not None: + sender.shutdown() + receiver = getattr(self, "_receiver", None) + if receiver is not None: + receiver.shutdown() # Deregister NIXL memory before shutting down components, so that # pinned GPU memory is released and can be re-allocated (e.g. when # the KV cache manager is recreated after profiling). - if self._finalizer is not None: - self._finalizer() + finalizer = getattr(self, "_finalizer", None) + if finalizer is not None: + finalizer() def __del__(self): try: self.shutdown() - except Exception: - pass + except Exception as e: + logger.warning(f"TransferWorker.__del__: exception during shutdown: {e}") def __enter__(self): return self - def __exit__(self, exc_type, exc_val, exc_tb): + def __exit__(self, _exc_type, _exc_val, _exc_tb): self.shutdown() diff --git a/tests/unittest/disaggregated/test_kv_transfer.py b/tests/unittest/disaggregated/test_kv_transfer.py index e603edc20c78..6347db2d6c09 100644 --- a/tests/unittest/disaggregated/test_kv_transfer.py +++ b/tests/unittest/disaggregated/test_kv_transfer.py @@ -17,7 +17,6 @@ from tensorrt_llm._torch.disaggregation.base.transfer import ( KVSlice, LayerRange, - SessionState, SessionStatus, TokenRange, ) @@ -124,26 +123,14 @@ def test_session_status_enum(): "INIT", "READY", "TRANSFERRING", - "TRANSFERRED", - "AUX_TRANSFERRED", - "COMPLETED", - "CANCELED", + "KV_TRANSFERRED", + "FULLY_TRANSFERRED", "ERROR", ] for name in expected: assert hasattr(SessionStatus, name) assert SessionStatus[name].value == name - assert len(SessionStatus) == 8 - - -def test_session_state_construction(): - state = SessionState(status=SessionStatus.INIT, finished_tasks=[]) - assert state.status == SessionStatus.INIT - assert state.finished_tasks == [] - - state2 = SessionState(status=SessionStatus.COMPLETED, finished_tasks=[1, 2, 3]) - assert state2.status == SessionStatus.COMPLETED - assert state2.finished_tasks == [1, 2, 3] + assert len(SessionStatus) == 6 def create_transfer_worker_setup( @@ -720,13 +707,13 @@ def add_and_verify_request( KVSlice(is_last_slice=True, block_ids_per_layer_groups=ctx_block_ids_per_group) for ctx_block_ids_per_group in ctx_block_ids_per_groups ] - send_slice_tasks = [ - sender_session._kv_tasks[sender_session.send(send_kv_slice)] + send_slice_futures = [ + sender_session.send(send_kv_slice) for sender_session, send_kv_slice in zip(sender_sessions, send_kv_slices) ] for sender_session in sender_sessions: - assert sender_session.state.status == SessionStatus.INIT + assert sender_session.status == SessionStatus.INIT receiver_sessions = [ gen_transfer_worker.create_rx_session(gen_request) @@ -736,8 +723,8 @@ def add_and_verify_request( KVSlice(is_last_slice=True, block_ids_per_layer_groups=gen_block_ids_per_group) for gen_block_ids_per_group in gen_block_ids_per_groups ] - recv_slice_tasks = [ - receiver_session._kv_tasks[receiver_session.receive(recv_kv_slice)] + recv_slice_futures = [ + receiver_session.receive(recv_kv_slice) for receiver_session, recv_kv_slice in zip(receiver_sessions, recv_kv_slices) ] @@ -750,8 +737,8 @@ def add_and_verify_request( KVSlice(is_last_slice=True, block_ids_per_layer_groups=gen_block_ids_per_group) for gen_block_ids_per_group in gen_block_ids_per_groups ] - recv_slice_tasks = [ - receiver_session._kv_tasks[receiver_session.receive(recv_kv_slice)] + recv_slice_futures = [ + receiver_session.receive(recv_kv_slice) for receiver_session, recv_kv_slice in zip(receiver_sessions, recv_kv_slices) ] @@ -765,37 +752,39 @@ def add_and_verify_request( time.sleep(0.1) for sender_session in sender_sessions: - assert sender_session.state.status != SessionStatus.INIT + assert sender_session.status != SessionStatus.INIT send_kv_slices = [ KVSlice(is_last_slice=True, block_ids_per_layer_groups=ctx_block_ids_per_group) for ctx_block_ids_per_group in ctx_block_ids_per_groups ] - send_slice_tasks = [ - sender_session._kv_tasks[sender_session.send(send_kv_slice)] + send_slice_futures = [ + sender_session.send(send_kv_slice) for sender_session, send_kv_slice in zip(sender_sessions, send_kv_slices) ] send_aux_tasks = [] for sender_session in sender_sessions: - sender_session.pack_aux() + sender_session.pack_aux(ctx_request) send_aux_tasks.append(sender_session.send_aux()) - for send_slice_task in send_slice_tasks: - send_slice_task.future.result() - for recv_slice_task in recv_slice_tasks: - recv_slice_task.future.result() + for future in send_slice_futures: + future.result() + for future in recv_slice_futures: + future.result() if not send_first: for send_aux_task in send_aux_tasks: send_aux_task.future.result() - sync_session_status = SessionStatus.TRANSFERRED if send_first else SessionStatus.AUX_TRANSFERRED + sync_session_status = ( + SessionStatus.KV_TRANSFERRED if send_first else SessionStatus.FULLY_TRANSFERRED + ) for sender_session in sender_sessions: - assert sender_session.state.status == sync_session_status + assert sender_session.status == sync_session_status if not send_first: time.sleep(0.1) for receiver_session in receiver_sessions: - assert receiver_session.state.status == sync_session_status, ( - f"receiver_session.state.status={receiver_session.state.status}, " + assert receiver_session.status == sync_session_status, ( + f"receiver_session.status={receiver_session.status}, " f"sync_session_status={sync_session_status} send_first={send_first}" ) @@ -912,9 +901,9 @@ def get_layers_in_group_per_pp(kv_cache_managers, pp_size, tp_size, group_id, is 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() + recv_session.unpack_aux(gen_request) - assert gen_request.context_phase_params.first_gen_tokens == [8 + ctx_request_id] + assert gen_request.py_first_gen_tokens == [8 + ctx_request_id] assert gen_request.py_draft_tokens == [ 9 + ctx_request_id, 10 + ctx_request_id, diff --git a/tests/unittest/disaggregated/test_kv_transfer_mp.py b/tests/unittest/disaggregated/test_kv_transfer_mp.py index 02f37a6f11f2..55455b815d85 100644 --- a/tests/unittest/disaggregated/test_kv_transfer_mp.py +++ b/tests/unittest/disaggregated/test_kv_transfer_mp.py @@ -332,11 +332,11 @@ def process_and_verify_request( # Get block ids and send block_ids = kv_cache_manager.get_batch_cache_indices([ctx_request.py_request_id])[0] send_kv_slice = KVSlice(is_last_slice=True, block_ids_per_layer_groups=[block_ids]) - send_slice_task = sender_session._kv_tasks[sender_session.send(send_kv_slice)] + send_future = sender_session.send(send_kv_slice) # Wait for send to complete - send_slice_task.future.result() - assert sender_session.state.status == SessionStatus.TRANSFERRED + send_future.result() + assert sender_session.status == SessionStatus.KV_TRANSFERRED # Get block data for verification block_data = kv_cache_manager.get_unique_primary_pool()[block_ids] @@ -371,11 +371,11 @@ def process_and_verify_request( # Get block ids and receive block_ids = kv_cache_manager.get_batch_cache_indices([gen_request.py_request_id])[0] recv_kv_slice = KVSlice(is_last_slice=True, block_ids_per_layer_groups=[block_ids]) - recv_slice_task = receiver_session._kv_tasks[receiver_session.receive(recv_kv_slice)] + recv_future = receiver_session.receive(recv_kv_slice) # Wait for receive to complete - recv_slice_task.future.result() - assert receiver_session.state.status == SessionStatus.TRANSFERRED + recv_future.result() + assert receiver_session.status == SessionStatus.KV_TRANSFERRED # Get block data for verification block_data = kv_cache_manager.get_unique_primary_pool()[block_ids] From cf973a26b3a5d4cc95a78044db5bab4dc1b5a839 Mon Sep 17 00:00:00 2001 From: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> Date: Fri, 20 Mar 2026 02:15:55 +0000 Subject: [PATCH 2/2] fix according to comments Signed-off-by: Shixiaowei02 <39303645+Shixiaowei02@users.noreply.github.com> --- .../_torch/disaggregation/native/transfer.py | 13 ++++++------- 1 file changed, 6 insertions(+), 7 deletions(-) diff --git a/tensorrt_llm/_torch/disaggregation/native/transfer.py b/tensorrt_llm/_torch/disaggregation/native/transfer.py index f005509cce0b..df6234641193 100644 --- a/tensorrt_llm/_torch/disaggregation/native/transfer.py +++ b/tensorrt_llm/_torch/disaggregation/native/transfer.py @@ -356,10 +356,10 @@ def _deliver_kv_to_agent(self, write_meta: WriteMeta): task.status = TaskStatus.TRANSFERRING agent_result = AgentResult.SUCCESS - if timer: - timer.record_transfer_start(write_meta.peer_rank) if write_meta.src_ptrs: request = Sender._make_agent_request(write_meta, device_id=self._device_id) + if timer: + timer.record_transfer_start(write_meta.peer_rank) if not self._agent.submit_transfer_requests(request).wait(): agent_result = AgentResult.FAILED if not write_meta.task_future.done(): @@ -673,10 +673,9 @@ def _save_peer_req_info(self, peer_transfer_req_info: RecvReqInfo): 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) if self._is_req_ready(req_info.unique_rid, expected_transfers): - if req_info.unique_rid in self._sessions: - session = self._get_session(req_info.unique_rid) - if session is not None and not session.receiver_ready: - session.receiver_ready = True + session = self._get_session(req_info.unique_rid) + if session is not None and not session.receiver_ready: + session.receiver_ready = True def has_all_peer_req_infos(self, unique_rid: int) -> bool: req_info = self._get_first_req_info(unique_rid) @@ -1073,7 +1072,7 @@ def _process_kv_agent_result(self, _send_id: bytes, message: list[bytes]): slice_id = int(slice_id_str) if msg_type.encode("ascii") != MessageType.KV_AGENT_RESULT: logger.error( - f"_process_kv_agent_result: unexpected msg_type={msg_type!r}, expected TASK_STATUS" + f"_process_kv_agent_result: unexpected msg_type={msg_type!r}, expected KV_AGENT_RESULT" ) return session = self._get_session(unique_rid)