From 2acb20f53ba5c741382537cf9107a0e306918ba0 Mon Sep 17 00:00:00 2001 From: Shangming Cai Date: Wed, 18 Mar 2026 08:26:18 +0800 Subject: [PATCH] [Disagg] Non-blocking `try_ensure_parallel_info` in pending queue, consolidate rank mapping into `PrefillServerInfo` (#20785) Signed-off-by: Shangming Cai Co-authored-by: hnyls2002 --- .../sglang/srt/disaggregation/common/conn.py | 238 +++++++++--------- python/sglang/srt/disaggregation/decode.py | 103 +++++--- 2 files changed, 178 insertions(+), 163 deletions(-) diff --git a/python/sglang/srt/disaggregation/common/conn.py b/python/sglang/srt/disaggregation/common/conn.py index 82832a429..0d72416da 100644 --- a/python/sglang/srt/disaggregation/common/conn.py +++ b/python/sglang/srt/disaggregation/common/conn.py @@ -46,6 +46,7 @@ logger = logging.getLogger(__name__) @dataclasses.dataclass class PrefillServerInfo: + # Topology fields (fetched from bootstrap server) attn_tp_size: int attn_cp_size: int dp_size: int @@ -54,6 +55,14 @@ class PrefillServerInfo: kv_cache_dtype: Optional[str] follow_bootstrap_room: bool + # Pre-computed rank mapping (set by try_ensure_parallel_info on decode side) + target_tp_rank: Optional[int] = None + target_tp_ranks: Optional[List[int]] = None + target_cp_ranks: Optional[List[int]] = None + target_pp_ranks: Optional[List[int]] = None + required_dst_info_num: Optional[int] = None + required_prefill_response_num: Optional[int] = None + def __post_init__(self): self.attn_tp_size = int(self.attn_tp_size) self.attn_cp_size = int(self.attn_cp_size) @@ -183,29 +192,29 @@ class CommonKVManager(BaseKVManager): with self.failure_lock: self.failure_records[bootstrap_room] = failure_reason - def ensure_parallel_info( - self, bootstrap_addr: str, max_retries: int = 5, retry_interval: float = 1.0 - ) -> bool: - """Fetch and cache prefill parallel info if not yet available. - Returns True if info is available (cached or freshly fetched). - Retries with backoff if the prefill server hasn't registered yet. - """ + def try_ensure_parallel_info(self, bootstrap_addr: str) -> bool: + """Single non-blocking attempt to fetch and cache prefill parallel info. + Returns True if info is available (cached or freshly fetched).""" if bootstrap_addr in self.prefill_info_table: return True - info = None - for attempt in range(max_retries): - info = self._fetch_prefill_server_info(bootstrap_addr) - if info is not None: - break - if attempt < max_retries - 1: - logger.info( - f"Prefill server info not available from {bootstrap_addr}, " - f"retrying ({attempt + 1}/{max_retries})..." + + info: PrefillServerInfo = None + try: + url = f"http://{bootstrap_addr}/route?prefill_dp_rank={-1}&prefill_cp_rank={-1}&target_tp_rank={-1}&target_pp_rank={-1}" + response = requests.get(url, timeout=5) + if response.status_code == 200: + data = response.json() + info = PrefillServerInfo(**data) + else: + logger.error( + f"Failed to get prefill server info: {response.status_code}, {response.text}" ) - time.sleep(retry_interval) - if info is None: + return False + except Exception as e: + logger.error(f"Error fetching prefill server info from bootstrap: {e}") return False + # Sanity checks if info.page_size is not None and info.page_size != self.kv_args.page_size: raise RuntimeError( f"Page size mismatch: prefill server has page_size={info.page_size}, " @@ -223,29 +232,89 @@ class CommonKVManager(BaseKVManager): f"Both servers must use the same --kv-cache-dtype value." ) + self._resolve_rank_mapping(info) self.prefill_info_table[bootstrap_addr] = info logger.debug(f"Prefill parallel info for [{bootstrap_addr}]: {info}") return True - @staticmethod - def _fetch_prefill_server_info( - bootstrap_addr: str, - ) -> Optional[PrefillServerInfo]: - """Fetch the prefill server info from the bootstrap server.""" - try: - url = f"http://{bootstrap_addr}/route?prefill_dp_rank={-1}&prefill_cp_rank={-1}&target_tp_rank={-1}&target_pp_rank={-1}" - response = requests.get(url, timeout=5) - if response.status_code == 200: - data = response.json() - return PrefillServerInfo(**data) - else: - logger.error( - f"Failed to get prefill server info: {response.status_code}, {response.text}" + def _resolve_rank_mapping(self, info: PrefillServerInfo) -> None: + """Compute TP/CP/PP rank mapping and store on the PrefillServerInfo object. + Deterministic for a given (bootstrap_addr, decode engine) pair.""" + # TP rank mapping + if self.attn_tp_size == info.attn_tp_size: + target_tp_rank = self.kv_args.engine_rank % self.attn_tp_size + required_dst_info_num = 1 + required_prefill_response_num = 1 + target_tp_ranks = [target_tp_rank] + elif self.attn_tp_size > info.attn_tp_size: + if not self.is_mla_backend: + logger.warning_once( + "Performance is NOT guaranteed when using different TP sizes for non-MLA models. " ) - return None - except Exception as e: - logger.error(f"Error fetching prefill server info from bootstrap: {e}") - return None + target_tp_rank = (self.kv_args.engine_rank % self.attn_tp_size) // ( + self.attn_tp_size // info.attn_tp_size + ) + required_dst_info_num = self.attn_tp_size // info.attn_tp_size + required_prefill_response_num = 1 + target_tp_ranks = [target_tp_rank] + else: + if not self.is_mla_backend: + logger.warning_once( + "Performance is NOT guaranteed when using different TP sizes for non-MLA models. " + ) + # For non-MLA models, one decode rank needs to retrieve KVCache from multiple prefill ranks + target_tp_ranks = list( + range( + (self.kv_args.engine_rank % self.attn_tp_size) + * (info.attn_tp_size // self.attn_tp_size), + (self.kv_args.engine_rank % self.attn_tp_size + 1) + * (info.attn_tp_size // self.attn_tp_size), + ) + ) + # For MLA models, we can retrieve KVCache from only one prefill rank, but we still need to maintain + # multiple connections in the connection pool and have to send dummy requests to other prefill ranks, + # or the KVPoll will never be set correctly + target_tp_rank = target_tp_ranks[0] + required_dst_info_num = 1 + if self.is_mla_backend: + required_prefill_response_num = 1 + else: + required_prefill_response_num = info.attn_tp_size // self.attn_tp_size + + # CP rank mapping — decode cp size should be equal to 1 + assert self.attn_cp_size == 1, ( + f"Decode cp size ({self.attn_cp_size}) should be equal to 1", + ) + if self.attn_cp_size == info.attn_cp_size: + assert info.attn_cp_size == 1, ( + f"When prefill cp size is 1, attn cp size should be 1, but got {self.attn_cp_size}", + ) + target_cp_ranks = [self.attn_cp_rank] + else: + target_cp_ranks = list(range(info.attn_cp_size)) + if not self.enable_all_cp_ranks_for_transfer: + # Only retrieve from prefill CP rank 0 when not using all ranks + target_cp_ranks = target_cp_ranks[:1] + required_prefill_response_num *= 1 + else: + required_prefill_response_num *= info.attn_cp_size // self.attn_cp_size + + # PP rank mapping — decode pp size should be equal to prefill pp size or 1 + assert self.pp_size == info.pp_size or self.pp_size == 1, ( + f"Decode pp size ({self.pp_size}) should be equal to prefill pp size ({info.pp_size}) or 1", + ) + if info.pp_size == self.pp_size: + target_pp_ranks = [self.pp_rank] + else: + target_pp_ranks = list(range(info.pp_size)) + required_prefill_response_num *= info.pp_size // self.pp_size + + info.target_tp_rank = target_tp_rank + info.target_tp_ranks = target_tp_ranks + info.target_cp_ranks = target_cp_ranks + info.target_pp_ranks = target_pp_ranks + info.required_dst_info_num = required_dst_info_num + info.required_prefill_response_num = required_prefill_response_num def register_to_bootstrap(self): """Register prefill server info to bootstrap server via HTTP POST.""" @@ -427,102 +496,25 @@ class CommonKVReceiver(BaseKVReceiver): self.kv_mgr = mgr self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Bootstrapping) - if not self.kv_mgr.ensure_parallel_info(self.bootstrap_addr): + if self.bootstrap_addr not in self.kv_mgr.prefill_info_table: self.kv_mgr.record_failure( self.bootstrap_room, - f"Could not fetch prefill parallel info from bootstrap_addr: {self.bootstrap_addr}", + f"Prefill server with bootstrap_addr: {self.bootstrap_addr} is healthy before, but now it is down. Request (bootstrap_room: {self.bootstrap_room}) has been marked as failed.", ) self.kv_mgr.update_status(self.bootstrap_room, KVPoll.Failed) self.bootstrap_infos = None return + # Read pre-computed rank mapping from prefill_info (computed in try_ensure_parallel_info) self.prefill_info = self.kv_mgr.prefill_info_table[self.bootstrap_addr] - - # Rank mapping for PD with different TP sizes per rank for target DP/CP group - if self.kv_mgr.attn_tp_size == self.prefill_info.attn_tp_size: - self.target_tp_rank = ( - self.kv_mgr.kv_args.engine_rank % self.kv_mgr.attn_tp_size - ) - self.required_dst_info_num = 1 - self.required_prefill_response_num = 1 - self.target_tp_ranks = [self.target_tp_rank] - elif self.kv_mgr.attn_tp_size > self.prefill_info.attn_tp_size: - if not self.kv_mgr.is_mla_backend: - logger.warning_once( - "Performance is NOT guaranteed when using different TP sizes for non-MLA models. " - ) - self.target_tp_rank = ( - self.kv_mgr.kv_args.engine_rank % self.kv_mgr.attn_tp_size - ) // (self.kv_mgr.attn_tp_size // self.prefill_info.attn_tp_size) - self.required_dst_info_num = ( - self.kv_mgr.attn_tp_size // self.prefill_info.attn_tp_size - ) - self.required_prefill_response_num = 1 - self.target_tp_ranks = [self.target_tp_rank] - else: - if not self.kv_mgr.is_mla_backend: - logger.warning_once( - "Performance is NOT guaranteed when using different TP sizes for non-MLA models. " - ) - # For non-MLA models, one decode rank needs to retrieve KVCache from multiple prefill ranks for non MLA models; - self.target_tp_ranks = [ - rank - for rank in range( - (self.kv_mgr.kv_args.engine_rank % self.kv_mgr.attn_tp_size) - * (self.prefill_info.attn_tp_size // self.kv_mgr.attn_tp_size), - (self.kv_mgr.kv_args.engine_rank % self.kv_mgr.attn_tp_size + 1) - * (self.prefill_info.attn_tp_size // self.kv_mgr.attn_tp_size), - ) - ] - - # For MLA models, we can retrieve KVCache from only one prefill rank, but we still need to maintain - # multiple connections in the connection pool and have to send dummy requests to other prefill ranks, - # or the KVPoll will never be set correctly - self.target_tp_rank = self.target_tp_ranks[0] - self.required_dst_info_num = 1 - if self.kv_mgr.is_mla_backend: - self.required_prefill_response_num = 1 - else: - self.required_prefill_response_num = ( - self.prefill_info.attn_tp_size // self.kv_mgr.attn_tp_size - ) - - # Decode cp size should be equal to 1 - assert self.kv_mgr.attn_cp_size == 1, ( - f"Decode cp size ({self.kv_mgr.attn_cp_size}) should be equal to 1", + self.target_tp_rank = self.prefill_info.target_tp_rank + self.target_tp_ranks = self.prefill_info.target_tp_ranks + self.target_cp_ranks = self.prefill_info.target_cp_ranks + self.target_pp_ranks = self.prefill_info.target_pp_ranks + self.required_dst_info_num = self.prefill_info.required_dst_info_num + self.required_prefill_response_num = ( + self.prefill_info.required_prefill_response_num ) - if self.kv_mgr.attn_cp_size == self.prefill_info.attn_cp_size: - # This means that the prefill cp size is 1 - assert self.prefill_info.attn_cp_size == 1, ( - f"When prefill cp size is 1, attn cp size should be 1, but got {self.kv_mgr.attn_cp_size}", - ) - self.target_cp_ranks = [self.kv_mgr.attn_cp_rank] - else: - self.target_cp_ranks = [ - rank for rank in range(self.prefill_info.attn_cp_size) - ] - if not self.kv_mgr.enable_all_cp_ranks_for_transfer: - # Only retrieve from prefill CP rank 0 when not using all ranks - self.target_cp_ranks = self.target_cp_ranks[:1] - self.required_prefill_response_num *= 1 - else: - self.required_prefill_response_num *= ( - self.prefill_info.attn_cp_size // self.kv_mgr.attn_cp_size - ) - - # Decode pp size should be equal to prefill pp size or 1 - assert ( - self.kv_mgr.pp_size == self.prefill_info.pp_size or self.kv_mgr.pp_size == 1 - ), ( - f"Decode pp size ({self.kv_mgr.pp_size}) should be equal to prefill pp size ({self.prefill_info.pp_size}) or 1", - ) - if self.prefill_info.pp_size == self.kv_mgr.pp_size: - self.target_pp_ranks = [self.kv_mgr.pp_rank] - else: - self.target_pp_ranks = [rank for rank in range(self.prefill_info.pp_size)] - self.required_prefill_response_num *= ( - self.prefill_info.pp_size // self.kv_mgr.pp_size - ) self.kv_mgr.required_prefill_response_num_table[self.bootstrap_room] = ( self.required_prefill_response_num diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 9d55ba027..0ede52313 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -84,6 +84,11 @@ def _is_fake_transfer(req: Req, server_args: ServerArgs) -> bool: ) +def _bootstrap_addr(req: Req) -> str: + # FIXME: make a property of a req + return f"{req.bootstrap_host}:{req.bootstrap_port}" + + class DecodeReqToTokenPool: """ The difference of DecodeReqToTokenPool and ReqToTokenPool is that @@ -265,6 +270,8 @@ class DecodePreallocQueue: self.queue: List[DecodeRequest] = [] self.retracted_queue: List[Req] = [] self.pending_reqs: List[Req] = [] + self._ensure_retry_count: Dict[str, int] = {} + self._max_ensure_retries: int = 30 # scheduling cycles self.kv_manager = self._init_kv_manager() if self.scheduler.tp_worker.is_hybrid_swa: @@ -352,22 +359,23 @@ class DecodePreallocQueue: req.retraction_mb_id = None self.retracted_queue.append(req) else: - prefill_dp_rank = self._resolve_prefill_dp_rank(req) - if prefill_dp_rank is None: - self.pending_reqs.append(req) + # NOTE: fake transfer does not need to resolve prefill dp rank in the pending queue + if _is_fake_transfer(req, self.scheduler.server_args): + self._create_receiver_and_enqueue(req, 0) return - self._create_receiver_and_enqueue(req, prefill_dp_rank) + + # Fast path: cache-only lookup, no network calls + prefill_dp_rank = self._resolve_prefill_dp_rank(req) + if prefill_dp_rank is not None: + self._create_receiver_and_enqueue(req, prefill_dp_rank) + else: + self.pending_reqs.append(req) def _resolve_prefill_dp_rank(self, req: Req) -> Optional[int]: if req.disagg_prefill_dp_rank is not None: return req.disagg_prefill_dp_rank - if _is_fake_transfer(req, self.scheduler.server_args): - return 0 - - bootstrap_addr = f"{req.bootstrap_host}:{req.bootstrap_port}" - - prefill_info = self.kv_manager.prefill_info_table.get(bootstrap_addr) + prefill_info = self.kv_manager.prefill_info_table.get(_bootstrap_addr(req)) if prefill_info is None: return None @@ -389,7 +397,7 @@ class DecodePreallocQueue: kv_receiver = kv_receiver_class( mgr=self.kv_manager, - bootstrap_addr=f"{req.bootstrap_host}:{req.bootstrap_port}", + bootstrap_addr=_bootstrap_addr(req), bootstrap_room=req.bootstrap_room, prefill_dp_rank=prefill_dp_rank, ) @@ -493,6 +501,40 @@ class DecodePreallocQueue: else: raise ValueError(f"Unexpected poll case: {poll}") + def _ensure_prefill_info( + self, addr_to_reqs: Dict[str, List[Req]] + ) -> Tuple[Dict[str, List[Req]], List[Req]]: + """Non-blocking ensure parallel info for each addr. + Returns (ready_addrs, remaining_reqs).""" + ready: Dict[str, List[Req]] = {} + remaining: List[Req] = [] + + for bootstrap_addr, reqs in addr_to_reqs.items(): + if self.kv_manager.try_ensure_parallel_info(bootstrap_addr): + if bootstrap_addr in self._ensure_retry_count: + del self._ensure_retry_count[bootstrap_addr] + ready[bootstrap_addr] = reqs + continue + + count = self._ensure_retry_count.get(bootstrap_addr, 0) + 1 + self._ensure_retry_count[bootstrap_addr] = count + + if count >= self._max_ensure_retries: + error_msg = f"Could not fetch prefill parallel info from {bootstrap_addr} after {count} attempts" + logger.error(error_msg) + for req in reqs: + prepare_abort( + req, error_msg, status_code=HTTPStatus.INTERNAL_SERVER_ERROR + ) + if self.scheduler.enable_metrics: + self.scheduler.metrics_collector.increment_bootstrap_failed_reqs() + self.scheduler.stream_output([req], req.return_logprob) + del self._ensure_retry_count[bootstrap_addr] + else: + remaining.extend(reqs) + + return ready, remaining + def _resolve_pending_reqs(self) -> None: """Batch-resolve prefill_dp_ranks for pending requests and create receivers.""" if not self.pending_reqs: @@ -501,34 +543,17 @@ class DecodePreallocQueue: # Group pending requests by bootstrap_addr addr_to_reqs: Dict[str, List[Req]] = {} for req in self.pending_reqs: - addr = f"{req.bootstrap_host}:{req.bootstrap_port}" + addr = _bootstrap_addr(req) addr_to_reqs.setdefault(addr, []).append(req) + # Pass 1: ensure parallel info for each addr + ready_addrs, remaining = self._ensure_prefill_info(addr_to_reqs) + + # Pass 2: resolve dp rank for addrs whose info is available resolved = [] - remaining = [] - - for bootstrap_addr, reqs in addr_to_reqs.items(): - # If a request is following the bootstrap room, - # we need get the prefill info before resolving the prefill_dp_ranks - # which is a conflict with the lazy resolve logic in CommonKVReceiver, - # so we need to ensure the parallel info before resolving it. - if not self.kv_manager.ensure_parallel_info(bootstrap_addr): - error_message = f"Could not fetch prefill parallel info from bootstrap server {bootstrap_addr}" - logger.error(error_message) - for req in reqs: - prepare_abort( - req, - error_message, - status_code=HTTPStatus.INTERNAL_SERVER_ERROR, - ) - if self.scheduler.enable_metrics: - self.scheduler.metrics_collector.increment_bootstrap_failed_reqs() - self.scheduler.stream_output([req], req.return_logprob) - continue - - need_query = [] + for bootstrap_addr, reqs in ready_addrs.items(): + need_query: List[Req] = [] for req in reqs: - # NOTE: we need resolve it again because we may ensure the parallel info here prefill_dp_rank = self._resolve_prefill_dp_rank(req) if prefill_dp_rank is not None: resolved.append((req, prefill_dp_rank)) @@ -536,16 +561,14 @@ class DecodePreallocQueue: need_query.append(req) if need_query: - from sglang.srt.disaggregation.common.conn import CommonKVReceiver - rooms = [req.bootstrap_room for req in need_query] room_to_rank = CommonKVReceiver.query_prefill_dp_ranks( bootstrap_addr, rooms ) for req in need_query: - room_key = str(req.bootstrap_room) - if room_key in room_to_rank: - resolved.append((req, int(room_to_rank[room_key]))) + prefill_dp_rank = room_to_rank.get(str(req.bootstrap_room)) + if prefill_dp_rank is not None: + resolved.append((req, int(prefill_dp_rank))) else: remaining.append(req)