From aa215e55230bc9b0eb503b5a3886c015c2e5f526 Mon Sep 17 00:00:00 2001 From: YAMY <74099316+YAMY1234@users.noreply.github.com> Date: Wed, 19 Aug 2026 02:25:52 -0700 Subject: [PATCH] [PD] Overlap prefill DP-rank bootstrap queries (#35071) --- python/sglang/srt/disaggregation/decode.py | 84 ++++++++++++++++++- .../test_decode_queue_cleanup.py | 46 ++++++++++ 2 files changed, 127 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 1085e3415..0b2eb3c85 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -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) diff --git a/test/registered/unit/disaggregation/test_decode_queue_cleanup.py b/test/registered/unit/disaggregation/test_decode_queue_cleanup.py index a0d07f686..82c593c2d 100644 --- a/test/registered/unit/disaggregation/test_decode_queue_cleanup.py +++ b/test/registered/unit/disaggregation/test_decode_queue_cleanup.py @@ -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")