[PD] Fix the infinite loop in deocde resolve_pending_reqs (#20371)

Signed-off-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
Shangming Cai
2026-03-11 14:11:19 -07:00
committed by GitHub
parent ab4b863546
commit af4c28904d
2 changed files with 49 additions and 33 deletions
@@ -184,7 +184,7 @@ class CommonKVManager(BaseKVManager):
self.failure_records[bootstrap_room] = failure_reason self.failure_records[bootstrap_room] = failure_reason
def ensure_parallel_info( def ensure_parallel_info(
self, bootstrap_addr: str, max_retries: int = 20, retry_interval: float = 1.0 self, bootstrap_addr: str, max_retries: int = 5, retry_interval: float = 1.0
) -> bool: ) -> bool:
"""Fetch and cache prefill parallel info if not yet available. """Fetch and cache prefill parallel info if not yet available.
Returns True if info is available (cached or freshly fetched). Returns True if info is available (cached or freshly fetched).
+24 -8
View File
@@ -24,7 +24,7 @@ import logging
from collections import deque from collections import deque
from dataclasses import dataclass from dataclasses import dataclass
from http import HTTPStatus from http import HTTPStatus
from typing import TYPE_CHECKING, List, Optional, Tuple from typing import TYPE_CHECKING, Dict, List, Optional, Tuple
import torch import torch
from torch.distributed import ProcessGroup from torch.distributed import ProcessGroup
@@ -498,18 +498,36 @@ class DecodePreallocQueue:
if not self.pending_reqs: if not self.pending_reqs:
return return
bootstrap_addr = f"{self.pending_reqs[0].bootstrap_host}:{self.pending_reqs[0].bootstrap_port}" # 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_to_reqs.setdefault(addr, []).append(req)
resolved = []
remaining = []
for bootstrap_addr, reqs in addr_to_reqs.items():
# If a request is following the bootstrap room, # If a request is following the bootstrap room,
# we need get the prefill info before resolving the prefill_dp_ranks # we need get the prefill info before resolving the prefill_dp_ranks
# which is a conflict with the lazy resolve logic in CommonKVReceiver, # which is a conflict with the lazy resolve logic in CommonKVReceiver,
# so we need to ensure the parallel info before resolving it. # so we need to ensure the parallel info before resolving it.
if not self.kv_manager.ensure_parallel_info(bootstrap_addr): if not self.kv_manager.ensure_parallel_info(bootstrap_addr):
return 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
resolved = []
need_query = [] need_query = []
for req in self.pending_reqs: for req in reqs:
# NOTE: we need resolve it again because we may ensure the parallel info here # NOTE: we need resolve it again because we may ensure the parallel info here
prefill_dp_rank = self._resolve_prefill_dp_rank(req) prefill_dp_rank = self._resolve_prefill_dp_rank(req)
if prefill_dp_rank is not None: if prefill_dp_rank is not None:
@@ -524,16 +542,14 @@ class DecodePreallocQueue:
room_to_rank = CommonKVReceiver.query_prefill_dp_ranks( room_to_rank = CommonKVReceiver.query_prefill_dp_ranks(
bootstrap_addr, rooms bootstrap_addr, rooms
) )
remaining = []
for req in need_query: for req in need_query:
room_key = str(req.bootstrap_room) room_key = str(req.bootstrap_room)
if room_key in room_to_rank: if room_key in room_to_rank:
resolved.append((req, int(room_to_rank[room_key]))) resolved.append((req, int(room_to_rank[room_key])))
else: else:
remaining.append(req) remaining.append(req)
self.pending_reqs = remaining self.pending_reqs = remaining
else:
self.pending_reqs = []
for req, prefill_dp_rank in resolved: for req, prefill_dp_rank in resolved:
self._create_receiver_and_enqueue(req, prefill_dp_rank) self._create_receiver_and_enqueue(req, prefill_dp_rank)