Files
sglang/python/sglang/srt/disaggregation/base/conn.py
T
2025-04-15 19:29:31 +08:00

114 lines
2.4 KiB
Python

from abc import ABC, abstractmethod
from typing import Optional
import numpy as np
import numpy.typing as npt
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.server_args import ServerArgs
class KVArgs:
engine_rank: int
kv_data_ptrs: list[int]
kv_data_lens: list[int]
kv_item_lens: list[int]
aux_data_ptrs: list[int]
aux_data_lens: list[int]
aux_item_lens: list[int]
ib_device: str
gpu_id: int
class KVPoll:
Failed = 0
Bootstrapping = 1
WaitingForInput = 2
Transferring = 3
Success = 4
class BaseKVManager(ABC):
"""Base class for managing transfers states"""
@abstractmethod
def __init__(
self,
args: KVArgs,
disaggregation_mode: DisaggregationMode,
server_args: ServerArgs,
): ...
class BaseKVSender(ABC):
@abstractmethod
def __init__(
self, mgr: BaseKVManager, bootstrap_addr: str, bootstrap_room: int
): ...
@abstractmethod
def init(self, num_kv_indices: int, aux_index: Optional[int] = None):
"""
Notify the decoder server about the kv indices length and aux index
"""
...
@abstractmethod
def send(self, kv_indices: npt.NDArray[np.int64]):
"""
Send the kv cache at the given kv indices to the decoder server
"""
...
@abstractmethod
def poll(self) -> KVPoll:
"""
Check the status of the kv cache transfer
"""
...
@abstractmethod
def failure_exception(self):
"""
Raise an exception if the kv cache transfer fails
"""
...
class BaseKVReceiver(ABC):
@abstractmethod
def __init__(
self,
mgr: BaseKVManager,
bootstrap_addr: str,
bootstrap_room: Optional[int] = None,
): ...
@abstractmethod
def init(self, kv_indices: npt.NDArray[np.int64], aux_index: Optional[int] = None):
"""
Notify the prefill server about the kv indices and aux index
"""
...
@abstractmethod
def poll(self) -> KVPoll:
"""
Check the status of the kv cache transfer
"""
...
@abstractmethod
def failure_exception(self):
"""
Raise an exception if the kv cache transfer fails
"""
...
class BaseKVBootstrapServer(ABC):
@abstractmethod
def __init__(self, port: int): ...