Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
152 changes: 66 additions & 86 deletions tensorrt_llm/_torch/disaggregation/base/transfer.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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]:
Expand All @@ -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):
Comment thread
Shixiaowei02 marked this conversation as resolved.
"""
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.
"""
...
1 change: 1 addition & 0 deletions tensorrt_llm/_torch/disaggregation/native/messenger.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}")
Expand Down
36 changes: 35 additions & 1 deletion tensorrt_llm/_torch/disaggregation/native/peer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Comment thread
Shixiaowei02 marked this conversation as resolved.
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):
Expand Down Expand Up @@ -199,13 +202,16 @@ 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]
peer_lg = peer_pt.layer_groups[peer_lg_idx]
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 = (
Expand Down Expand Up @@ -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
Loading
Loading