[PD] Add transfer backend abstraction (#5328)

This commit is contained in:
Byron Hsu
2025-04-14 01:39:39 +08:00
committed by GitHub
parent f765579046
commit a9499885e9
12 changed files with 236 additions and 41 deletions
+18 -5
View File
@@ -28,10 +28,19 @@ import numpy as np
import torch
from torch.distributed import ProcessGroup
from sglang.srt.disaggregation.conn import KVArgs, KVManager, KVPoll, KVReceiver
from sglang.srt.disaggregation.base import (
BaseKVManager,
BaseKVReceiver,
BaseKVSender,
KVArgs,
KVPoll,
)
from sglang.srt.disaggregation.utils import (
DisaggregationMode,
KVClassType,
ReqToMetadataIdxAllocator,
TransferBackend,
get_kv_class,
poll_and_all_reduce,
)
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
@@ -51,7 +60,7 @@ if TYPE_CHECKING:
@dataclass
class DecodeRequest:
req: Req
kv_receiver: KVReceiver
kv_receiver: BaseKVReceiver
waiting_for_input: bool = False
metadata_buffer_index: int = -1
@@ -75,6 +84,7 @@ class DecodePreallocQueue:
tp_rank: int,
tp_size: int,
bootstrap_port: int,
transfer_backend: TransferBackend,
):
self.req_to_token_pool = req_to_token_pool
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
@@ -94,9 +104,10 @@ class DecodePreallocQueue:
# Queue for requests pending pre-allocation
self.queue: List[DecodeRequest] = []
self.transfer_backend = transfer_backend
self.kv_manager = self._init_kv_manager()
def _init_kv_manager(self) -> KVManager:
def _init_kv_manager(self) -> BaseKVManager:
kv_args = KVArgs()
kv_args.engine_rank = self.tp_rank
kv_data_ptrs, kv_data_lens, kv_item_lens = (
@@ -117,13 +128,15 @@ class DecodePreallocQueue:
metadata_buffer[0].nbytes for metadata_buffer in self.metadata_buffers
]
kv_args.ib_device = "mock-ib-device"
kv_manager = KVManager(kv_args, DisaggregationMode("decode"))
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
kv_manager = kv_manager_class(kv_args, DisaggregationMode.DECODE)
return kv_manager
def add(self, req: Req) -> None:
"""Add a request to the pending queue."""
kv_receiver = KVReceiver(
kv_receiver_class = get_kv_class(self.transfer_backend, KVClassType.RECEIVER)
kv_receiver = kv_receiver_class(
mgr=self.kv_manager,
bootstrap_addr=f"{req.bootstrap_host}:{self.bootstrap_port}",
bootstrap_room=req.bootstrap_room,