[PD] Fix the infinite loop in deocde resolve_pending_reqs (#20371)
Signed-off-by: Shangming Cai <csmthu@gmail.com>
This commit is contained in:
@@ -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,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,42 +498,58 @@ 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]] = {}
|
||||||
# If a request is following the bootstrap room,
|
for req in self.pending_reqs:
|
||||||
# we need get the prefill info before resolving the prefill_dp_ranks
|
addr = f"{req.bootstrap_host}:{req.bootstrap_port}"
|
||||||
# which is a conflict with the lazy resolve logic in CommonKVReceiver,
|
addr_to_reqs.setdefault(addr, []).append(req)
|
||||||
# so we need to ensure the parallel info before resolving it.
|
|
||||||
if not self.kv_manager.ensure_parallel_info(bootstrap_addr):
|
|
||||||
return
|
|
||||||
|
|
||||||
resolved = []
|
resolved = []
|
||||||
need_query = []
|
remaining = []
|
||||||
for req in self.pending_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))
|
|
||||||
else:
|
|
||||||
need_query.append(req)
|
|
||||||
|
|
||||||
if need_query:
|
for bootstrap_addr, reqs in addr_to_reqs.items():
|
||||||
from sglang.srt.disaggregation.common.conn import CommonKVReceiver
|
# 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
|
||||||
|
|
||||||
rooms = [req.bootstrap_room for req in need_query]
|
need_query = []
|
||||||
room_to_rank = CommonKVReceiver.query_prefill_dp_ranks(
|
for req in reqs:
|
||||||
bootstrap_addr, rooms
|
# NOTE: we need resolve it again because we may ensure the parallel info here
|
||||||
)
|
prefill_dp_rank = self._resolve_prefill_dp_rank(req)
|
||||||
remaining = []
|
if prefill_dp_rank is not None:
|
||||||
for req in need_query:
|
resolved.append((req, prefill_dp_rank))
|
||||||
room_key = str(req.bootstrap_room)
|
|
||||||
if room_key in room_to_rank:
|
|
||||||
resolved.append((req, int(room_to_rank[room_key])))
|
|
||||||
else:
|
else:
|
||||||
remaining.append(req)
|
need_query.append(req)
|
||||||
self.pending_reqs = remaining
|
|
||||||
else:
|
if need_query:
|
||||||
self.pending_reqs = []
|
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])))
|
||||||
|
else:
|
||||||
|
remaining.append(req)
|
||||||
|
|
||||||
|
self.pending_reqs = remaining
|
||||||
|
|
||||||
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)
|
||||||
|
|||||||
Reference in New Issue
Block a user