[PD] Overlap prefill DP-rank bootstrap queries (#35071)

This commit is contained in:
YAMY
2026-08-19 17:25:52 +08:00
committed by GitHub
parent 0e4a09480c
commit aa215e5523
2 changed files with 127 additions and 3 deletions
+81 -3
View File
@@ -23,6 +23,7 @@ from __future__ import annotations
import logging
import time
from collections import deque
from concurrent.futures import Future
from dataclasses import dataclass
from http import HTTPStatus
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
@@ -345,6 +346,10 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
self.queue: List[DecodeRequest] = []
self.retracted_queue: List[Req] = []
self.pending_reqs: List[DecodeRequest] = []
# In-flight authoritative room -> DP-rank lookups, consumed below.
self._prefill_dp_rank_queries: Dict[
str, Tuple[Tuple[int, ...], Future[Dict[str, int]]]
] = {}
self._ensure_retry_count: Dict[str, int] = {}
self._max_ensure_retries: int = 15 # scheduling cycles
self._ensure_last_attempt_time: Dict[str, float] = {}
@@ -694,6 +699,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
self.add(req, is_retracted=is_retracted)
def release_memory_occupation(self):
self._cancel_prefill_dp_rank_queries()
self.queue.clear()
for req in self.retracted_queue:
retraction_discard(
@@ -874,9 +880,52 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
return ready, remaining
def prefetch_prefill_dp_rank_queries(self) -> None:
"""Start DP-rank lookups before their normal consume point."""
if not self.pending_reqs:
return
queries = self._prefill_dp_rank_queries
addr_to_reqs: Dict[str, List[DecodeRequest]] = {}
for decode_req in self.pending_reqs:
addr = _bootstrap_addr(decode_req.req)
addr_to_reqs.setdefault(addr, []).append(decode_req)
for bootstrap_addr in set(queries) - set(addr_to_reqs):
_, stale_future = queries.pop(bootstrap_addr)
stale_future.cancel()
for bootstrap_addr, decode_reqs in addr_to_reqs.items():
if bootstrap_addr in queries:
continue
if self.kv_manager.prefill_info_table.get(bootstrap_addr) is None:
continue
rooms = tuple(
decode_req.req.bootstrap_room
for decode_req in decode_reqs
if self._resolve_prefill_dp_rank(decode_req.req) is None
)
if not rooms:
continue
future = self.kv_manager._ensure_prefill_recompute_executor().submit(
CommonKVReceiver.query_prefill_dp_ranks,
bootstrap_addr,
list(rooms),
)
queries[bootstrap_addr] = (rooms, future)
def _cancel_prefill_dp_rank_queries(self) -> None:
for _, future in self._prefill_dp_rank_queries.values():
future.cancel()
self._prefill_dp_rank_queries.clear()
def _resolve_pending_reqs(self) -> None:
"""Batch-resolve prefill_dp_ranks for pending requests and initialize receivers."""
if not self.pending_reqs:
self._cancel_prefill_dp_rank_queries()
return
# Group pending requests by bootstrap_addr
@@ -901,9 +950,26 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
# Pass 2: resolve dp rank for addrs whose info is available
if need_query:
rooms = [decode_req.req.bootstrap_room for decode_req in need_query]
room_to_rank = CommonKVReceiver.query_prefill_dp_ranks(
bootstrap_addr, rooms
)
prefetched = self._prefill_dp_rank_queries.pop(bootstrap_addr, None)
prefetched_rooms = prefetched[0] if prefetched is not None else ()
if (
prefetched is not None
and tuple(rooms[: len(prefetched_rooms)]) == prefetched_rooms
):
room_to_rank = prefetched[1].result()
remaining_rooms = rooms[len(prefetched_rooms) :]
if remaining_rooms:
room_to_rank.update(
CommonKVReceiver.query_prefill_dp_ranks(
bootstrap_addr, remaining_rooms
)
)
else:
if prefetched is not None:
prefetched[1].cancel()
room_to_rank = CommonKVReceiver.query_prefill_dp_ranks(
bootstrap_addr, rooms
)
for decode_req in need_query:
prefill_dp_rank = room_to_rank.get(
str(decode_req.req.bootstrap_room)
@@ -912,6 +978,10 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
resolved.append((decode_req, int(prefill_dp_rank)))
else:
remaining.append(decode_req)
else:
prefetched = self._prefill_dp_rank_queries.pop(bootstrap_addr, None)
if prefetched is not None:
prefetched[1].cancel()
self.pending_reqs = remaining
@@ -2232,6 +2302,10 @@ class SchedulerDisaggregationDecodeMixin:
"""A normal scheduler loop for decode worker in disaggregation mode."""
while True:
# Pending rooms from the prior cycle can overlap request intake and
# the tail of the in-flight decode graph.
if not self._engine_paused:
self.disagg_decode_prealloc_queue.prefetch_prefill_dp_rank_queries()
# Receive requests
recv_reqs = self.request_receiver.recv_requests()
self.process_input_requests(recv_reqs)
@@ -2271,6 +2345,10 @@ class SchedulerDisaggregationDecodeMixin:
self.process_batch_result(tmp_batch, tmp_result)
while True:
# Pending rooms from the prior cycle can overlap request intake and
# the tail of the in-flight decode graph.
if not self._engine_paused:
self.disagg_decode_prealloc_queue.prefetch_prefill_dp_rank_queries()
# Receive requests
recv_reqs = self.request_receiver.recv_requests()
self.process_input_requests(recv_reqs)