[PD-Disagg] Improve type hints across all conn.py (#19208)

This commit is contained in:
Liangsheng Yin
2026-02-23 19:44:23 -08:00
committed by GitHub
parent edba96b98a
commit feb041f4e5
6 changed files with 72 additions and 44 deletions

View File

@@ -281,7 +281,7 @@ class CommonKVManager(BaseKVManager):
class CommonKVSender(BaseKVSender):
def __init__(
self,
mgr: BaseKVManager,
mgr: CommonKVManager,
bootstrap_addr: str,
bootstrap_room: int,
dest_tp_ranks: List[int],
@@ -342,7 +342,7 @@ class CommonKVReceiver(BaseKVReceiver):
def __init__(
self,
mgr: BaseKVManager,
mgr: CommonKVManager,
bootstrap_addr: str,
bootstrap_room: Optional[int] = None,
prefill_dp_rank: Optional[int] = None,

View File

@@ -25,14 +25,15 @@ import time
from collections import deque
from dataclasses import dataclass
from http import HTTPStatus
from typing import TYPE_CHECKING, List, Optional, Tuple, Type
from typing import TYPE_CHECKING, List, Optional, Tuple
import torch
from torch.distributed import ProcessGroup
from sglang.srt.configs.mamba_utils import Mamba2CacheParams
from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE
from sglang.srt.disaggregation.base import BaseKVManager, BaseKVReceiver, KVPoll
from sglang.srt.disaggregation.base import KVPoll
from sglang.srt.disaggregation.common.conn import CommonKVManager, CommonKVReceiver
from sglang.srt.disaggregation.utils import (
FAKE_BOOTSTRAP_HOST,
DisaggregationMode,
@@ -69,10 +70,18 @@ logger = logging.getLogger(__name__)
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.server_args import ServerArgs
CLIP_MAX_NEW_TOKEN = get_int_env_var("SGLANG_CLIP_MAX_NEW_TOKENS_ESTIMATION", 4096)
def _is_fake_transfer(req: Req, server_args: ServerArgs) -> bool:
return req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
req.bootstrap_host is None
and server_args.disaggregation_transfer_backend == "fake"
)
class DecodeReqToTokenPool:
"""
The difference of DecodeReqToTokenPool and ReqToTokenPool is that
@@ -193,7 +202,7 @@ class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool):
@dataclass
class DecodeRequest:
req: Req
kv_receiver: BaseKVReceiver
kv_receiver: CommonKVReceiver
waiting_for_input: bool = False
metadata_buffer_index: int = -1
@@ -264,7 +273,7 @@ class DecodePreallocQueue:
self.scheduler.tp_worker.model_runner.swa_max_total_num_tokens,
)
def _init_kv_manager(self) -> BaseKVManager:
def _init_kv_manager(self) -> CommonKVManager:
kv_args_class = get_kv_class(self.transfer_backend, KVClassType.KVARGS)
kv_args = kv_args_class()
@@ -326,10 +335,8 @@ class DecodePreallocQueue:
kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device
kv_args.gpu_id = self.scheduler.gpu_id
kv_manager_class: Type[BaseKVManager] = get_kv_class(
self.transfer_backend, KVClassType.MANAGER
)
kv_manager: BaseKVManager = kv_manager_class(
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
kv_manager = kv_manager_class(
kv_args,
DisaggregationMode.DECODE,
self.scheduler.server_args,
@@ -356,10 +363,7 @@ class DecodePreallocQueue:
if req.data_parallel_rank is not None:
return req.data_parallel_rank
if req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
req.bootstrap_host is None
and self.scheduler.server_args.disaggregation_transfer_backend == "fake"
):
if _is_fake_transfer(req, self.scheduler.server_args):
return 0
bootstrap_addr = f"{req.bootstrap_host}:{req.bootstrap_port}"
@@ -374,15 +378,12 @@ class DecodePreallocQueue:
return None
def _create_receiver_and_enqueue(self, req: Req, dp_rank: int) -> None:
if req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
req.bootstrap_host is None
and self.scheduler.server_args.disaggregation_transfer_backend == "fake"
):
kv_receiver_class = get_kv_class(TransferBackend.FAKE, KVClassType.RECEIVER)
else:
kv_receiver_class = get_kv_class(
self.transfer_backend, KVClassType.RECEIVER
)
backend = (
TransferBackend.FAKE
if _is_fake_transfer(req, self.scheduler.server_args)
else self.transfer_backend
)
kv_receiver_class = get_kv_class(backend, KVClassType.RECEIVER)
kv_receiver = kv_receiver_class(
mgr=self.kv_manager,
@@ -833,11 +834,7 @@ class DecodeTransferQueue:
else 0
)
if decode_req.req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
decode_req.req.bootstrap_host is None
and self.scheduler.server_args.disaggregation_transfer_backend == "fake"
):
# Warm up or fake transfer mode
if _is_fake_transfer(decode_req.req, self.scheduler.server_args):
pass
elif actual_room == 0:
# Case 1: Metadata not ready yet (actual_room == 0)

View File

@@ -23,11 +23,12 @@ import logging
import time
from collections import deque
from http import HTTPStatus
from typing import TYPE_CHECKING, List, Optional, Type
from typing import TYPE_CHECKING, List, Optional
import torch
from sglang.srt.disaggregation.base import BaseKVManager, KVPoll
from sglang.srt.disaggregation.base import KVPoll
from sglang.srt.disaggregation.common.conn import CommonKVManager
from sglang.srt.disaggregation.utils import (
FAKE_BOOTSTRAP_HOST,
DisaggregationMode,
@@ -134,7 +135,7 @@ class PrefillBootstrapQueue:
self.scheduler.tp_worker.model_runner.swa_max_total_num_tokens,
)
def _init_kv_manager(self) -> BaseKVManager:
def _init_kv_manager(self) -> CommonKVManager:
kv_args_class = get_kv_class(self.transfer_backend, KVClassType.KVARGS)
kv_args = kv_args_class()
kv_args.engine_rank = self.tp_rank
@@ -197,10 +198,8 @@ class PrefillBootstrapQueue:
kv_args.state_item_lens = []
kv_args.state_type = "none"
kv_manager_class: Type[BaseKVManager] = get_kv_class(
self.transfer_backend, KVClassType.MANAGER
)
kv_manager: BaseKVManager = kv_manager_class(
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
kv_manager = kv_manager_class(
kv_args,
DisaggregationMode.PREFILL,
self.scheduler.server_args,
@@ -212,10 +211,12 @@ class PrefillBootstrapQueue:
if self._check_if_req_exceed_kv_capacity(req):
return
if req.bootstrap_host == FAKE_BOOTSTRAP_HOST:
kv_sender_class = get_kv_class(TransferBackend.FAKE, KVClassType.SENDER)
else:
kv_sender_class = get_kv_class(self.transfer_backend, KVClassType.SENDER)
backend = (
TransferBackend.FAKE
if req.bootstrap_host == FAKE_BOOTSTRAP_HOST
else self.transfer_backend
)
kv_sender_class = get_kv_class(backend, KVClassType.SENDER)
dest_tp_ranks = [self.tp_rank]

View File

@@ -5,7 +5,7 @@ import random
from collections import deque
from contextlib import nullcontext
from enum import Enum
from typing import TYPE_CHECKING, Optional, Type
from typing import TYPE_CHECKING, Literal, Optional, Type, overload
import numpy as np
import torch
@@ -15,6 +15,13 @@ from sglang.srt.environ import envs
from sglang.srt.utils import is_npu
if TYPE_CHECKING:
from sglang.srt.disaggregation.base.conn import KVArgs
from sglang.srt.disaggregation.common.conn import (
CommonKVBootstrapServer,
CommonKVManager,
CommonKVReceiver,
CommonKVSender,
)
from sglang.srt.managers.schedule_batch import Req
#########################
@@ -260,6 +267,28 @@ class KVClassType(Enum):
BOOTSTRAP_SERVER = "bootstrap_server"
@overload
def get_kv_class(
transfer_backend: TransferBackend, class_type: Literal[KVClassType.KVARGS]
) -> Type[KVArgs]: ...
@overload
def get_kv_class(
transfer_backend: TransferBackend, class_type: Literal[KVClassType.MANAGER]
) -> Type[CommonKVManager]: ...
@overload
def get_kv_class(
transfer_backend: TransferBackend, class_type: Literal[KVClassType.SENDER]
) -> Type[CommonKVSender]: ...
@overload
def get_kv_class(
transfer_backend: TransferBackend, class_type: Literal[KVClassType.RECEIVER]
) -> Type[CommonKVReceiver]: ...
@overload
def get_kv_class(
transfer_backend: TransferBackend, class_type: Literal[KVClassType.BOOTSTRAP_SERVER]
) -> Type[CommonKVBootstrapServer]: ...
def get_kv_class(
transfer_backend: TransferBackend, class_type: KVClassType
) -> Optional[Type]:

View File

@@ -1,9 +1,7 @@
"""Start bootstrap/kv-store-related server"""
import os
from typing import Type
from sglang.srt.disaggregation.base import BaseKVBootstrapServer
from sglang.srt.disaggregation.utils import (
DisaggregationMode,
KVClassType,
@@ -22,10 +20,10 @@ def start_disagg_service(
if disagg_mode == DisaggregationMode.PREFILL:
# only start bootstrap server on prefill tm
kv_bootstrap_server_class: Type[BaseKVBootstrapServer] = get_kv_class(
kv_bootstrap_server_class = get_kv_class(
transfer_backend, KVClassType.BOOTSTRAP_SERVER
)
bootstrap_server: BaseKVBootstrapServer = kv_bootstrap_server_class(
bootstrap_server = kv_bootstrap_server_class(
host=server_args.host,
port=server_args.disaggregation_bootstrap_port,
dp_size=server_args.dp_size,

View File

@@ -2551,6 +2551,9 @@ class ServerArgs:
logger.warning("KV cache is forced as chunk cache for decode server")
elif self.disaggregation_mode == "prefill":
assert (
self.disaggregation_transfer_backend != "fake"
), "Prefill server does not support 'fake' as the transfer backend"
if self.disaggregation_decode_tp is None:
self.disaggregation_decode_tp = self.tp_size
if self.disaggregation_decode_dp is None: