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
5 changes: 2 additions & 3 deletions tensorrt_llm/_torch/disaggregation/base/agent.py
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
import os
from abc import ABC, abstractmethod
from dataclasses import dataclass
from typing import List, NamedTuple, Optional
from typing import List, NamedTuple, Optional, Tuple

from tensorrt_llm import logger

Expand All @@ -25,7 +25,6 @@ class MemoryDesc(NamedTuple):
ptr: int
size: int
device_id: int
name: Optional[str] = None


@dataclass
Expand All @@ -46,7 +45,7 @@ class TransferRequest:
@dataclass
class RegMemoryDescs:
type: str
descs: List[MemoryDesc]
descs: List[Tuple[int, int, int, str]]


class TransferStatus(ABC):
Expand Down
78 changes: 25 additions & 53 deletions tensorrt_llm/_torch/disaggregation/base/transfer.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from enum import Enum
from typing import List, Optional
from typing import List, Optional, cast

from tensorrt_llm import DisaggregatedParams
from tensorrt_llm._torch.pyexecutor.llm_request import LlmRequest
Expand Down Expand Up @@ -108,74 +108,46 @@ class ReceiverBase(ABC):
...


class TxSessionBase(ABC):
def __init__(self, sender: SenderBase, args: SessionArgsBase):
"""
Initializes the transmission session.
:param sender: The sender instance responsible for sending data.
:param args: The session arguments.
"""
self._sender = sender
class _SessionBase(ABC):
"""Shared base for Tx/Rx sessions."""

def __init__(self, args: SessionArgsBase):
self._base_args = args

@property
def disagg_request_id(self) -> int:
return self._base_args.params.disagg_request_id
return cast(int, self._base_args.params.disagg_request_id)

@abstractmethod
def send(self, slice: KVSlice) -> concurrent.futures.Future:
"""
Sends a slice of KV cache data and returns a Future for the transfer.
:param slice: The KV slice to send.
"""
...
def is_completed(self) -> bool: ...

@property
@abstractmethod
def exception(self) -> Optional[Exception]:
"""
Returns any exception that occurred during the session.
"""
...
def wait_complete(self) -> Optional[WaitResult]: ...

@property
@abstractmethod
def close(self) -> None:
"""
Closes the session and releases any resources.
"""
...
def exception(self) -> Optional[Exception]: ...

@abstractmethod
def close(self) -> None: ...

class RxSessionBase(ABC):
def __init__(self, receiver: ReceiverBase, args: SessionArgsBase):
"""
Initializes the reception session.
:param receiver: The receiver instance responsible for receiving data.
"""
self._receiver = receiver
self._base_args = args

@property
def disagg_request_id(self) -> int:
return self._base_args.params.disagg_request_id
class TxSessionBase(_SessionBase):
def __init__(self, sender: SenderBase, args: SessionArgsBase):
super().__init__(args)
self._sender = sender

@abstractmethod
def receive(self, slice: KVSlice) -> concurrent.futures.Future:
"""
Receives a slice of KV cache data and returns a Future for the transfer.
:param slice: The KV slice to receive.
"""
...
def send(self, slice: KVSlice) -> concurrent.futures.Future: ...


class RxSessionBase(_SessionBase):
def __init__(self, receiver: ReceiverBase, args: SessionArgsBase):
super().__init__(args)
self._receiver = receiver

@property
@abstractmethod
def exception(self) -> Optional[Exception]:
"""Returns any exception that occurred during the session."""
...
def receive(self, slice: KVSlice) -> concurrent.futures.Future: ...

@abstractmethod
def close(self) -> None:
"""
Closes the session and releases any resources.
"""
...
def wait_complete(self, blocking: bool = False) -> Optional[WaitResult]: ...
Loading
Loading