[PD] Add transfer backend abstraction (#5328)
This commit is contained in:
@@ -0,0 +1,8 @@
|
||||
from .conn import (
|
||||
BaseKVBootstrapServer,
|
||||
BaseKVManager,
|
||||
BaseKVReceiver,
|
||||
BaseKVSender,
|
||||
KVArgs,
|
||||
KVPoll,
|
||||
)
|
||||
@@ -0,0 +1,106 @@
|
||||
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
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
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): ...
|
||||
|
||||
|
||||
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): ...
|
||||
Reference in New Issue
Block a user