[PD] Overlap prefill DP-rank bootstrap queries (#35071)
This commit is contained in:
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user