Signed-off-by: Shangming Cai <caishangming@linux.alibaba.com> Co-authored-by: ybyang <ybyang7@iflytek.com>
114 lines
2.4 KiB
Python
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): ...
|