[PD] Overlap prefill DP-rank bootstrap queries (#35071)
This commit is contained in:
@@ -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,9 +950,26 @@ 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]
|
||||||
room_to_rank = CommonKVReceiver.query_prefill_dp_ranks(
|
prefetched = self._prefill_dp_rank_queries.pop(bootstrap_addr, None)
|
||||||
bootstrap_addr, rooms
|
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:
|
for decode_req in need_query:
|
||||||
prefill_dp_rank = room_to_rank.get(
|
prefill_dp_rank = room_to_rank.get(
|
||||||
str(decode_req.req.bootstrap_room)
|
str(decode_req.req.bootstrap_room)
|
||||||
@@ -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")
|
||||||
|
|||||||
Reference in New Issue
Block a user