[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)
@@ -1,4 +1,5 @@
import unittest
from concurrent.futures import Future
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
@@ -189,6 +190,51 @@ class TestDecodeQueueCleanup(CustomTestCase):
self.assertEqual(ready, {})
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.prepare_abort")
@patch("sglang.srt.disaggregation.decode.poll_and_all_reduce")