[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)
|
||||
|
||||
Reference in New Issue
Block a user