[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
@@ -23,6 +23,7 @@ from __future__ import annotations
import logging import logging
import time import time
from collections import deque from collections import deque
from concurrent.futures import Future
from dataclasses import dataclass from dataclasses import dataclass
from http import HTTPStatus from http import HTTPStatus
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple
@@ -345,6 +346,10 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
self.queue: List[DecodeRequest] = [] self.queue: List[DecodeRequest] = []
self.retracted_queue: List[Req] = [] self.retracted_queue: List[Req] = []
self.pending_reqs: List[DecodeRequest] = [] 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._ensure_retry_count: Dict[str, int] = {}
self._max_ensure_retries: int = 15 # scheduling cycles self._max_ensure_retries: int = 15 # scheduling cycles
self._ensure_last_attempt_time: Dict[str, float] = {} self._ensure_last_attempt_time: Dict[str, float] = {}
@@ -694,6 +699,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
self.add(req, is_retracted=is_retracted) self.add(req, is_retracted=is_retracted)
def release_memory_occupation(self): def release_memory_occupation(self):
self._cancel_prefill_dp_rank_queries()
self.queue.clear() self.queue.clear()
for req in self.retracted_queue: for req in self.retracted_queue:
retraction_discard( retraction_discard(
@@ -874,9 +880,52 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
return ready, remaining 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: def _resolve_pending_reqs(self) -> None:
"""Batch-resolve prefill_dp_ranks for pending requests and initialize receivers.""" """Batch-resolve prefill_dp_ranks for pending requests and initialize receivers."""
if not self.pending_reqs: if not self.pending_reqs:
self._cancel_prefill_dp_rank_queries()
return return
# Group pending requests by bootstrap_addr # Group pending requests by bootstrap_addr
@@ -901,6 +950,23 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
# Pass 2: resolve dp rank for addrs whose info is available # Pass 2: resolve dp rank for addrs whose info is available
if need_query: if need_query:
rooms = [decode_req.req.bootstrap_room for decode_req in need_query] rooms = [decode_req.req.bootstrap_room for decode_req in need_query]
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( room_to_rank = CommonKVReceiver.query_prefill_dp_ranks(
bootstrap_addr, rooms bootstrap_addr, rooms
) )
@@ -912,6 +978,10 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
resolved.append((decode_req, int(prefill_dp_rank))) resolved.append((decode_req, int(prefill_dp_rank)))
else: else:
remaining.append(decode_req) 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 self.pending_reqs = remaining
@@ -2232,6 +2302,10 @@ class SchedulerDisaggregationDecodeMixin:
"""A normal scheduler loop for decode worker in disaggregation mode.""" """A normal scheduler loop for decode worker in disaggregation mode."""
while True: 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 # Receive requests
recv_reqs = self.request_receiver.recv_requests() recv_reqs = self.request_receiver.recv_requests()
self.process_input_requests(recv_reqs) self.process_input_requests(recv_reqs)
@@ -2271,6 +2345,10 @@ class SchedulerDisaggregationDecodeMixin:
self.process_batch_result(tmp_batch, tmp_result) self.process_batch_result(tmp_batch, tmp_result)
while True: 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 # Receive requests
recv_reqs = self.request_receiver.recv_requests() recv_reqs = self.request_receiver.recv_requests()
self.process_input_requests(recv_reqs) self.process_input_requests(recv_reqs)
@@ -1,4 +1,5 @@
import unittest import unittest
from concurrent.futures import Future
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
@@ -189,6 +190,51 @@ class TestDecodeQueueCleanup(CustomTestCase):
self.assertEqual(ready, {}) self.assertEqual(ready, {})
self.assertEqual(remaining, []) self.assertEqual(remaining, [])
def test_prefetches_prefill_dp_rank_query(self):
addr = "127.0.0.1:11500"
executor = MagicMock()
future = Future()
future.set_result({"7": 1})
executor.submit.return_value = future
def decode_req(room):
return SimpleNamespace(
req=SimpleNamespace(
bootstrap_host="127.0.0.1",
bootstrap_port=11500,
bootstrap_room=room,
),
kv_receiver=MagicMock(),
)
first = decode_req(7)
queue = DecodePreallocQueue.__new__(DecodePreallocQueue)
queue.pending_reqs = [first]
queue._prefill_dp_rank_queries = {}
queue.kv_manager = SimpleNamespace(
prefill_info_table={addr: object()},
_ensure_prefill_recompute_executor=lambda: executor,
)
queue._resolve_prefill_dp_rank = MagicMock(return_value=None)
queue._ensure_prefill_info = lambda groups: (groups, [])
queue.prefetch_prefill_dp_rank_queries()
tail = decode_req(8)
queue.pending_reqs.append(tail)
with patch(
"sglang.srt.disaggregation.decode."
"CommonKVReceiver.query_prefill_dp_ranks",
return_value={"8": 2},
) as query:
queue._resolve_pending_reqs()
_, called_addr, called_rooms = executor.submit.call_args.args
self.assertEqual((called_addr, called_rooms), (addr, [7]))
query.assert_called_once_with(addr, [8])
first.kv_receiver.init.assert_called_once_with(1)
tail.kv_receiver.init.assert_called_once_with(2)
self.assertEqual(queue.pending_reqs, [])
@patch("sglang.srt.disaggregation.decode.release_kv_cache") @patch("sglang.srt.disaggregation.decode.release_kv_cache")
@patch("sglang.srt.disaggregation.decode.prepare_abort") @patch("sglang.srt.disaggregation.decode.prepare_abort")
@patch("sglang.srt.disaggregation.decode.poll_and_all_reduce") @patch("sglang.srt.disaggregation.decode.poll_and_all_reduce")