[PDD] Add true request retraction for PDD (#25372)
Signed-off-by: Ata Fatahi <immrata@gmail.com>
This commit is contained in:
@@ -1,6 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import concurrent.futures
|
||||||
import dataclasses
|
import dataclasses
|
||||||
import logging
|
import logging
|
||||||
import threading
|
import threading
|
||||||
@@ -76,6 +77,12 @@ class PrefillServerInfo:
|
|||||||
kv_cache_dtype: Optional[str]
|
kv_cache_dtype: Optional[str]
|
||||||
follow_bootstrap_room: bool
|
follow_bootstrap_room: bool
|
||||||
|
|
||||||
|
# PD true-retraction rebootstrap: the prefill's HTTP API port. The decode
|
||||||
|
# already knows the prefill host (the bootstrap_addr host), so it can POST
|
||||||
|
# /generate to http://{bootstrap_host}:{prefill_http_port} to trigger a KV
|
||||||
|
# recompute -- no router-injected pd_rebootstrap_prefill_url needed.
|
||||||
|
prefill_http_port: Optional[int] = None
|
||||||
|
|
||||||
# Pre-computed rank mapping (set by try_ensure_parallel_info on decode side)
|
# Pre-computed rank mapping (set by try_ensure_parallel_info on decode side)
|
||||||
target_tp_rank: Optional[int] = None
|
target_tp_rank: Optional[int] = None
|
||||||
target_tp_ranks: Optional[List[int]] = None
|
target_tp_ranks: Optional[List[int]] = None
|
||||||
@@ -94,6 +101,9 @@ class PrefillServerInfo:
|
|||||||
str(self.kv_cache_dtype) if self.kv_cache_dtype is not None else None
|
str(self.kv_cache_dtype) if self.kv_cache_dtype is not None else None
|
||||||
)
|
)
|
||||||
self.follow_bootstrap_room = bool(self.follow_bootstrap_room)
|
self.follow_bootstrap_room = bool(self.follow_bootstrap_room)
|
||||||
|
self.prefill_http_port = (
|
||||||
|
int(self.prefill_http_port) if self.prefill_http_port is not None else None
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@dataclasses.dataclass
|
@dataclasses.dataclass
|
||||||
@@ -204,6 +214,16 @@ class CommonKVManager(BaseKVManager):
|
|||||||
# fail to receive the KV Cache transfer done signal after bootstrapping.
|
# fail to receive the KV Cache transfer done signal after bootstrapping.
|
||||||
# These timeout requests should be aborted to release the tree cache.
|
# These timeout requests should be aborted to release the tree cache.
|
||||||
self.waiting_timeout = envs.SGLANG_DISAGGREGATION_WAITING_TIMEOUT.get()
|
self.waiting_timeout = envs.SGLANG_DISAGGREGATION_WAITING_TIMEOUT.get()
|
||||||
|
# PD true-retraction rebootstrap: a shared executor + per-thread HTTP
|
||||||
|
# sessions used to drive the original prefill worker's ``/generate``
|
||||||
|
# endpoint so it recomputes a retracted request's prefix KV under the
|
||||||
|
# current weights. Created lazily on first use so deployments that
|
||||||
|
# never retract pay nothing.
|
||||||
|
self._prefill_recompute_executor: Optional[
|
||||||
|
concurrent.futures.ThreadPoolExecutor
|
||||||
|
] = None
|
||||||
|
self._prefill_recompute_executor_lock = threading.Lock()
|
||||||
|
self._prefill_recompute_sessions = threading.local()
|
||||||
else:
|
else:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Unsupported DisaggregationMode: {self.disaggregation_mode}"
|
f"Unsupported DisaggregationMode: {self.disaggregation_mode}"
|
||||||
@@ -233,6 +253,151 @@ class CommonKVManager(BaseKVManager):
|
|||||||
with self.failure_lock:
|
with self.failure_lock:
|
||||||
self.failure_records[bootstrap_room] = failure_reason
|
self.failure_records[bootstrap_room] = failure_reason
|
||||||
|
|
||||||
|
def _ensure_prefill_recompute_executor(
|
||||||
|
self,
|
||||||
|
) -> concurrent.futures.ThreadPoolExecutor:
|
||||||
|
"""Lazily create the shared executor that drives PD-retract rebootstrap
|
||||||
|
``/generate`` calls. One executor per (decode) kv manager, shared across
|
||||||
|
all receivers."""
|
||||||
|
executor = self._prefill_recompute_executor
|
||||||
|
if executor is not None:
|
||||||
|
return executor
|
||||||
|
with self._prefill_recompute_executor_lock:
|
||||||
|
if self._prefill_recompute_executor is None:
|
||||||
|
workers = envs.SGLANG_DISAGGREGATION_THREAD_POOL_SIZE.get()
|
||||||
|
if workers is None:
|
||||||
|
workers = 16
|
||||||
|
self._prefill_recompute_executor = (
|
||||||
|
concurrent.futures.ThreadPoolExecutor(
|
||||||
|
max_workers=max(1, workers),
|
||||||
|
thread_name_prefix="pd-rebootstrap-prefill",
|
||||||
|
)
|
||||||
|
)
|
||||||
|
return self._prefill_recompute_executor
|
||||||
|
|
||||||
|
def _get_prefill_recompute_session(self) -> requests.Session:
|
||||||
|
"""Per-thread ``requests.Session`` for the rebootstrap executor threads
|
||||||
|
(``requests.Session`` is not safe for concurrent cross-thread use)."""
|
||||||
|
session = getattr(self._prefill_recompute_sessions, "session", None)
|
||||||
|
if session is None:
|
||||||
|
session = requests.Session()
|
||||||
|
self._prefill_recompute_sessions.session = session
|
||||||
|
return session
|
||||||
|
|
||||||
|
def _resolve_rebootstrap_prefill_url(
|
||||||
|
self, kv_receiver: CommonKVReceiver
|
||||||
|
) -> Optional[str]:
|
||||||
|
"""Derive the prefill ``/generate`` base URL for a PD true-retraction
|
||||||
|
rebootstrap from bootstrap info.
|
||||||
|
|
||||||
|
The decode already knows the prefill host (the ``bootstrap_addr`` host,
|
||||||
|
which the router/client set to the prefill's HTTP host), and the prefill
|
||||||
|
self-registers its HTTP API port in ``PrefillServerInfo`` at bootstrap
|
||||||
|
registration. Combining them yields ``http://{host}:{prefill_http_port}``
|
||||||
|
with no router-injected ``pd_rebootstrap_prefill_url``.
|
||||||
|
"""
|
||||||
|
prefill_info = self.prefill_info_table.get(kv_receiver.bootstrap_addr)
|
||||||
|
if prefill_info is None or prefill_info.prefill_http_port is None:
|
||||||
|
return None
|
||||||
|
host = NetworkAddress.parse(kv_receiver.bootstrap_addr).host
|
||||||
|
return NetworkAddress(host, prefill_info.prefill_http_port).to_url()
|
||||||
|
|
||||||
|
def submit_prefill_recompute(
|
||||||
|
self, kv_receiver: CommonKVReceiver, payload: dict
|
||||||
|
) -> None:
|
||||||
|
"""Dispatch a PD true-retraction rebootstrap ``/generate`` to the
|
||||||
|
original prefill worker so it recomputes the retracted request's prefix
|
||||||
|
KV under the current weights and transfers it back over the
|
||||||
|
already-bootstrapped channel.
|
||||||
|
|
||||||
|
The target prefill ``/generate`` URL is derived from bootstrap info (the
|
||||||
|
prefill self-registered its HTTP port), not from a router-injected field.
|
||||||
|
|
||||||
|
Non-blocking from the scheduler's perspective: the HTTP POST runs on the
|
||||||
|
shared executor. Any failure (unresolved URL, HTTP error, exception) is
|
||||||
|
surfaced through the standard ``KVPoll.Failed`` path via
|
||||||
|
``kv_receiver.abort()`` so the scheduler's existing transfer-failure
|
||||||
|
handling streams the aborted request back to the client. ``payload`` is
|
||||||
|
prebuilt by the decode scheduler (``Req.build_rebootstrap_payload``) so
|
||||||
|
HTTP/sampling concerns stay on the kv manager.
|
||||||
|
|
||||||
|
The decode scheduler broadcasts each retracted request to every rank in
|
||||||
|
its attention TP/CP group and every PP stage, so all of them reach this
|
||||||
|
call and would each POST an identical ``/generate`` -- making the prefill
|
||||||
|
worker recompute the same request once per decode rank. The ``/generate``
|
||||||
|
is a server-level call: the prefill frontend fans it out to its own
|
||||||
|
workers and transfers the recomputed KV back to *all* decode ranks, so
|
||||||
|
exactly one decode rank must issue it. Elect the same leader the request
|
||||||
|
receiver uses (attn-tp/attn-cp group leader, first PP stage); the other
|
||||||
|
ranks still bootstrap and receive their KV shard as usual, and on failure
|
||||||
|
the leader-only abort matches the leader-only output streaming (other
|
||||||
|
ranks fall back to the per-request waiting-timeout safety net).
|
||||||
|
"""
|
||||||
|
if self.attn_tp_rank != 0 or self.attn_cp_rank != 0 or self.pp_rank != 0:
|
||||||
|
return
|
||||||
|
prefill_url = self._resolve_rebootstrap_prefill_url(kv_receiver)
|
||||||
|
if not prefill_url:
|
||||||
|
logger.error(
|
||||||
|
"PD retract rebootstrap could not resolve the prefill /generate "
|
||||||
|
"URL from bootstrap info (rid=%s bootstrap_room=%s bootstrap_addr=%s).",
|
||||||
|
payload.get("rid"),
|
||||||
|
payload.get("bootstrap_room"),
|
||||||
|
kv_receiver.bootstrap_addr,
|
||||||
|
)
|
||||||
|
self._fail_prefill_recompute(
|
||||||
|
kv_receiver,
|
||||||
|
"PD retract rebootstrap could not resolve the prefill /generate "
|
||||||
|
"URL from bootstrap info.",
|
||||||
|
)
|
||||||
|
return
|
||||||
|
self._ensure_prefill_recompute_executor().submit(
|
||||||
|
self._run_prefill_recompute, kv_receiver, prefill_url, payload
|
||||||
|
)
|
||||||
|
|
||||||
|
def _fail_prefill_recompute(
|
||||||
|
self, kv_receiver: CommonKVReceiver, reason: str
|
||||||
|
) -> None:
|
||||||
|
"""Fail a rebootstrap request via the standard ``KVPoll.Failed`` path.
|
||||||
|
|
||||||
|
``abort()`` transitions the receiver to Failed and notifies the prefill
|
||||||
|
worker to release its orphaned bootstrap entry, but records a generic
|
||||||
|
reason; we overwrite it with a descriptive one so the eventual
|
||||||
|
``failure_exception`` (and the client-facing abort message) explains that
|
||||||
|
the rebootstrap ``/generate`` failed rather than reporting a spurious
|
||||||
|
``AbortReq``.
|
||||||
|
"""
|
||||||
|
kv_receiver.abort()
|
||||||
|
self.record_failure(kv_receiver.bootstrap_room, reason)
|
||||||
|
|
||||||
|
def _run_prefill_recompute(
|
||||||
|
self, kv_receiver: CommonKVReceiver, prefill_url: str, payload: dict
|
||||||
|
) -> None:
|
||||||
|
rid = payload.get("rid")
|
||||||
|
try:
|
||||||
|
response = self._get_prefill_recompute_session().post(
|
||||||
|
prefill_url.rstrip("/") + "/generate",
|
||||||
|
json=payload,
|
||||||
|
timeout=self.waiting_timeout,
|
||||||
|
)
|
||||||
|
if response.status_code >= 400:
|
||||||
|
logger.error(
|
||||||
|
"PD rebootstrap prefill failed for rid=%s status=%s body=%s",
|
||||||
|
rid,
|
||||||
|
response.status_code,
|
||||||
|
response.text[:512],
|
||||||
|
)
|
||||||
|
self._fail_prefill_recompute(
|
||||||
|
kv_receiver,
|
||||||
|
f"PD retract rebootstrap /generate failed for rid={rid} "
|
||||||
|
f"(status={response.status_code}).",
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.exception("PD rebootstrap prefill request failed for rid=%s", rid)
|
||||||
|
self._fail_prefill_recompute(
|
||||||
|
kv_receiver,
|
||||||
|
f"PD retract rebootstrap /generate request errored for rid={rid}.",
|
||||||
|
)
|
||||||
|
|
||||||
def try_ensure_parallel_info(self, bootstrap_addr: str) -> bool:
|
def try_ensure_parallel_info(self, bootstrap_addr: str) -> bool:
|
||||||
"""Single non-blocking attempt to fetch and cache prefill parallel info.
|
"""Single non-blocking attempt to fetch and cache prefill parallel info.
|
||||||
Returns True if info is available (cached or freshly fetched)."""
|
Returns True if info is available (cached or freshly fetched)."""
|
||||||
@@ -420,6 +585,10 @@ class CommonKVManager(BaseKVManager):
|
|||||||
"page_size": self.kv_args.page_size,
|
"page_size": self.kv_args.page_size,
|
||||||
"kv_cache_dtype": self.server_args.kv_cache_dtype,
|
"kv_cache_dtype": self.server_args.kv_cache_dtype,
|
||||||
"load_balance_method": self.server_args.load_balance_method,
|
"load_balance_method": self.server_args.load_balance_method,
|
||||||
|
# Self-register the HTTP API port so the decode can derive the PD
|
||||||
|
# retract rebootstrap /generate URL from bootstrap info instead of a
|
||||||
|
# router-injected pd_rebootstrap_prefill_url.
|
||||||
|
"prefill_http_port": self.server_args.port,
|
||||||
}
|
}
|
||||||
|
|
||||||
max_retries, initial_delay, max_delay = 5, 1.0, 30.0
|
max_retries, initial_delay, max_delay = 5, 1.0, 30.0
|
||||||
@@ -594,10 +763,9 @@ class CommonKVManager(BaseKVManager):
|
|||||||
"""
|
"""
|
||||||
start_layer = self.kv_args.prefill_start_layer
|
start_layer = self.kv_args.prefill_start_layer
|
||||||
end_layer = getattr(self.kv_args, "prefill_end_layer", None)
|
end_layer = getattr(self.kv_args, "prefill_end_layer", None)
|
||||||
assert end_layer is not None, (
|
assert (
|
||||||
"KVArgs.prefill_end_layer must be set when using "
|
end_layer is not None
|
||||||
"compressed-MLA PD with PP"
|
), "KVArgs.prefill_end_layer must be set when using compressed-MLA PD with PP"
|
||||||
)
|
|
||||||
|
|
||||||
c4_full = sum(1 for r in mla_ratios if r == 4)
|
c4_full = sum(1 for r in mla_ratios if r == 4)
|
||||||
c128_full = sum(1 for r in mla_ratios if r == 128)
|
c128_full = sum(1 for r in mla_ratios if r == 128)
|
||||||
@@ -1231,6 +1399,7 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
|
|||||||
self.page_size = None
|
self.page_size = None
|
||||||
self.kv_cache_dtype: Optional[str] = None
|
self.kv_cache_dtype: Optional[str] = None
|
||||||
self.follow_bootstrap_room: Optional[bool] = None
|
self.follow_bootstrap_room: Optional[bool] = None
|
||||||
|
self.prefill_http_port: Optional[int] = None
|
||||||
self.prefill_port_table: Dict[
|
self.prefill_port_table: Dict[
|
||||||
int, Dict[int, Dict[int, Dict[int, PrefillRankInfo]]]
|
int, Dict[int, Dict[int, Dict[int, PrefillRankInfo]]]
|
||||||
] = {}
|
] = {}
|
||||||
@@ -1297,6 +1466,7 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
|
|||||||
rank_port = int(data["rank_port"])
|
rank_port = int(data["rank_port"])
|
||||||
page_size = int(data["page_size"])
|
page_size = int(data["page_size"])
|
||||||
kv_cache_dtype = data["kv_cache_dtype"]
|
kv_cache_dtype = data["kv_cache_dtype"]
|
||||||
|
prefill_http_port = data.get("prefill_http_port")
|
||||||
|
|
||||||
if self.attn_tp_size is None:
|
if self.attn_tp_size is None:
|
||||||
self.attn_tp_size = attn_tp_size
|
self.attn_tp_size = attn_tp_size
|
||||||
@@ -1316,6 +1486,9 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
|
|||||||
if self.kv_cache_dtype is None and kv_cache_dtype is not None:
|
if self.kv_cache_dtype is None and kv_cache_dtype is not None:
|
||||||
self.kv_cache_dtype = kv_cache_dtype
|
self.kv_cache_dtype = kv_cache_dtype
|
||||||
|
|
||||||
|
if self.prefill_http_port is None and prefill_http_port is not None:
|
||||||
|
self.prefill_http_port = int(prefill_http_port)
|
||||||
|
|
||||||
if self.follow_bootstrap_room is None:
|
if self.follow_bootstrap_room is None:
|
||||||
load_balance_method = data.get(
|
load_balance_method = data.get(
|
||||||
"load_balance_method", "follow_bootstrap_room"
|
"load_balance_method", "follow_bootstrap_room"
|
||||||
@@ -1385,6 +1558,7 @@ class CommonKVBootstrapServer(BaseKVBootstrapServer):
|
|||||||
if self.follow_bootstrap_room is not None
|
if self.follow_bootstrap_room is not None
|
||||||
else True
|
else True
|
||||||
),
|
),
|
||||||
|
prefill_http_port=self.prefill_http_port,
|
||||||
)
|
)
|
||||||
return web.json_response(dataclasses.asdict(info), status=200)
|
return web.json_response(dataclasses.asdict(info), status=200)
|
||||||
|
|
||||||
|
|||||||
@@ -255,6 +255,7 @@ class DecodeRequest:
|
|||||||
kv_receiver: CommonKVReceiver
|
kv_receiver: CommonKVReceiver
|
||||||
waiting_for_input: bool = False
|
waiting_for_input: bool = False
|
||||||
metadata_buffer_index: int = -1
|
metadata_buffer_index: int = -1
|
||||||
|
is_rebootstrap: bool = False
|
||||||
|
|
||||||
# HiCache Status
|
# HiCache Status
|
||||||
prefix_match: Optional[DecodePrefixMatch] = None
|
prefix_match: Optional[DecodePrefixMatch] = None
|
||||||
@@ -326,6 +327,13 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
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] = {}
|
||||||
self._ensure_retry_interval: float = 1.0 # seconds
|
self._ensure_retry_interval: float = 1.0 # seconds
|
||||||
|
# Retracted requests staged for rebootstrap while generation is paused.
|
||||||
|
# Enqueued into ``self.queue`` only on ``continue_generation`` so the
|
||||||
|
# prefix KV is recomputed under the post-retract (updated) weights.
|
||||||
|
# NOTE: requests held here are not reachable by ``/abort_request``; to
|
||||||
|
# support aborting them we would need an additional fix in the
|
||||||
|
# scheduler. In practice this shouldn't arise in the RL scenario.
|
||||||
|
self.held_rebootstrap_reqs: List[Req] = []
|
||||||
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
||||||
if self.enable_staging and self.is_mla_backend:
|
if self.enable_staging and self.is_mla_backend:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
@@ -373,7 +381,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
return self._swa_tail_len(len(req.origin_input_ids)) + len(req.output_ids)
|
return self._swa_tail_len(len(req.origin_input_ids)) + len(req.output_ids)
|
||||||
|
|
||||||
def _prealloc_kv_lens(self, req: Req) -> Tuple[int, int]:
|
def _prealloc_kv_lens(self, req: Req) -> Tuple[int, int]:
|
||||||
allocated_kv_len = len(req.origin_input_ids) + max(len(req.output_ids) - 1, 0)
|
allocated_kv_len = self._pre_alloc_fill_len(req)
|
||||||
if self._uses_swa_tail_prealloc():
|
if self._uses_swa_tail_prealloc():
|
||||||
return allocated_kv_len, self._swa_tail_len(allocated_kv_len)
|
return allocated_kv_len, self._swa_tail_len(allocated_kv_len)
|
||||||
return allocated_kv_len, allocated_kv_len
|
return allocated_kv_len, allocated_kv_len
|
||||||
@@ -477,8 +485,18 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
)
|
)
|
||||||
return kv_manager
|
return kv_manager
|
||||||
|
|
||||||
def add(self, req: Req, is_retracted: bool = False) -> None:
|
def add(
|
||||||
"""Add a request to the pending queue."""
|
self, req: Req, is_retracted: bool = False, is_rebootstrap: bool = False
|
||||||
|
) -> None:
|
||||||
|
"""Add a request to the pending queue.
|
||||||
|
|
||||||
|
``is_rebootstrap`` marks a PD true-retraction request whose prefix KV
|
||||||
|
must be recomputed by the original prefill worker under the current
|
||||||
|
weights (rather than resumed from stale CPU KV). It otherwise follows the
|
||||||
|
same bootstrap-handshake path as a fresh request; the ``/generate``
|
||||||
|
dispatch happens later, after preallocation and ``send_metadata`` (see
|
||||||
|
``pop_preallocated``).
|
||||||
|
"""
|
||||||
if self._check_if_req_exceed_kv_capacity(req):
|
if self._check_if_req_exceed_kv_capacity(req):
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -486,7 +504,9 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
req.retraction_mb_id = None
|
req.retraction_mb_id = None
|
||||||
self.retracted_queue.append(req)
|
self.retracted_queue.append(req)
|
||||||
else:
|
else:
|
||||||
decode_req = self._create_receiver_and_enqueue(req)
|
decode_req = self._create_receiver_and_enqueue(
|
||||||
|
req, is_rebootstrap=is_rebootstrap
|
||||||
|
)
|
||||||
|
|
||||||
# NOTE: fake transfer does not need to resolve prefill dp rank in the pending queue
|
# NOTE: fake transfer does not need to resolve prefill dp rank in the pending queue
|
||||||
if _is_fake_transfer(req, self.scheduler.server_args):
|
if _is_fake_transfer(req, self.scheduler.server_args):
|
||||||
@@ -538,7 +558,9 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
|
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _create_receiver_and_enqueue(self, req: Req) -> DecodeRequest:
|
def _create_receiver_and_enqueue(
|
||||||
|
self, req: Req, is_rebootstrap: bool = False
|
||||||
|
) -> DecodeRequest:
|
||||||
backend = (
|
backend = (
|
||||||
TransferBackend.FAKE
|
TransferBackend.FAKE
|
||||||
if _is_fake_transfer(req, self.scheduler.server_args)
|
if _is_fake_transfer(req, self.scheduler.server_args)
|
||||||
@@ -552,13 +574,59 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
bootstrap_room=req.bootstrap_room,
|
bootstrap_room=req.bootstrap_room,
|
||||||
)
|
)
|
||||||
|
|
||||||
decode_req = DecodeRequest(req=req, kv_receiver=kv_receiver)
|
decode_req = DecodeRequest(
|
||||||
|
req=req, kv_receiver=kv_receiver, is_rebootstrap=is_rebootstrap
|
||||||
|
)
|
||||||
self.queue.append(decode_req)
|
self.queue.append(decode_req)
|
||||||
return decode_req
|
return decode_req
|
||||||
|
|
||||||
|
def hold_rebootstrap(self, req: Req) -> None:
|
||||||
|
"""Stage a retracted request for rebootstrap without enqueuing it yet.
|
||||||
|
|
||||||
|
Retraction is always paired with a weight update
|
||||||
|
(``pause_generation(mode="retract")`` -> ``update_weights`` ->
|
||||||
|
``continue_generation``). Enqueuing the rebootstrap into ``self.queue``
|
||||||
|
here would leave the preallocation queue non-empty, which makes the
|
||||||
|
scheduler non-idle so ``update_weights``' post-update cache flush
|
||||||
|
asserts and crashes the decode worker. Instead we hold the request and
|
||||||
|
enqueue it from ``enqueue_held_rebootstrap`` on resume, so its prefix KV
|
||||||
|
is recomputed by the prefill worker under the updated weights.
|
||||||
|
"""
|
||||||
|
self.held_rebootstrap_reqs.append(req)
|
||||||
|
|
||||||
|
def enqueue_held_rebootstrap(self) -> None:
|
||||||
|
"""Enqueue all staged rebootstrap requests when generation resumes."""
|
||||||
|
held = self.held_rebootstrap_reqs
|
||||||
|
self.held_rebootstrap_reqs = []
|
||||||
|
for req in held:
|
||||||
|
self.add(req, is_rebootstrap=True)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _rebootstrap_prefill_len(req: Req) -> int:
|
||||||
|
if getattr(req, "pd_rebootstrap_in_progress", False):
|
||||||
|
return len(req.origin_input_ids) + len(req.output_ids)
|
||||||
|
return len(req.origin_input_ids)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _pre_alloc_fill_len(req: Req) -> int:
|
||||||
|
if getattr(req, "pd_rebootstrap_in_progress", False):
|
||||||
|
# pause_generation(retract) already popped the boundary token out of
|
||||||
|
# output_ids (it is replayed via the decode-side override at commit
|
||||||
|
# time), so output_ids here is prompt + emitted-tokens-minus-boundary,
|
||||||
|
# i.e. the original seqlen - 1. The prefill recomputes KV for *all* of
|
||||||
|
# these tokens, leaving no just-sampled "pending" token in the list, so
|
||||||
|
# we allocate exactly len(origin)+len(output_ids) with no -1 (unlike
|
||||||
|
# normal decode, where the last token's KV has not been written yet).
|
||||||
|
# This is the same token count as offloading-based retraction, where
|
||||||
|
# offload_kv_cache saves seqlen-1 tokens; the boundary token's KV is
|
||||||
|
# (re)computed on the decode side once generation resumes.
|
||||||
|
return len(req.origin_input_ids) + len(req.output_ids)
|
||||||
|
return len(req.origin_input_ids) + max(len(req.output_ids) - 1, 0)
|
||||||
|
|
||||||
def _check_if_req_exceed_kv_capacity(self, req: Req) -> bool:
|
def _check_if_req_exceed_kv_capacity(self, req: Req) -> bool:
|
||||||
if len(req.origin_input_ids) > self.max_total_num_tokens:
|
input_len = self._rebootstrap_prefill_len(req)
|
||||||
message = f"Request {req.rid} exceeds the maximum number of tokens: {len(req.origin_input_ids)} > {self.max_total_num_tokens}"
|
if input_len > self.max_total_num_tokens:
|
||||||
|
message = f"Request {req.rid} exceeds the maximum number of tokens: {input_len} > {self.max_total_num_tokens}"
|
||||||
logger.error(message)
|
logger.error(message)
|
||||||
prepare_abort(req, message, status_code=HTTPStatus.BAD_REQUEST)
|
prepare_abort(req, message, status_code=HTTPStatus.BAD_REQUEST)
|
||||||
self.scheduler.output_streamer.stream_output([req], req.return_logprob)
|
self.scheduler.output_streamer.stream_output([req], req.return_logprob)
|
||||||
@@ -830,10 +898,11 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
if rids_to_check is not None and decode_req.req.rid not in rids_to_check:
|
if rids_to_check is not None and decode_req.req.rid not in rids_to_check:
|
||||||
continue
|
continue
|
||||||
if isinstance(decode_req.req.finished_reason, FINISH_ABORT):
|
if isinstance(decode_req.req.finished_reason, FINISH_ABORT):
|
||||||
self.scheduler.output_streamer.stream_output(
|
if not getattr(decode_req.req, "finished_output", False):
|
||||||
[decode_req.req],
|
self.scheduler.output_streamer.stream_output(
|
||||||
decode_req.req.return_logprob,
|
[decode_req.req],
|
||||||
)
|
decode_req.req.return_logprob,
|
||||||
|
)
|
||||||
decode_req.kv_receiver.clear()
|
decode_req.kv_receiver.clear()
|
||||||
decode_req.kv_receiver = None
|
decode_req.kv_receiver = None
|
||||||
failed_reqs.append(decode_req)
|
failed_reqs.append(decode_req)
|
||||||
@@ -884,9 +953,13 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
|
|
||||||
# Memory estimation: don't add if the projected memory cannot be met
|
# Memory estimation: don't add if the projected memory cannot be met
|
||||||
# TODO: add new_token ratio
|
# TODO: add new_token ratio
|
||||||
origin_input_len = len(decode_req.req.origin_input_ids)
|
origin_input_len = self._rebootstrap_prefill_len(decode_req.req)
|
||||||
prefix_match: Optional[DecodePrefixMatch] = None
|
prefix_match: Optional[DecodePrefixMatch] = None
|
||||||
if self.scheduler.server_args.disaggregation_decode_enable_radix_cache:
|
use_decode_radix_cache = (
|
||||||
|
self.scheduler.server_args.disaggregation_decode_enable_radix_cache
|
||||||
|
and not decode_req.is_rebootstrap
|
||||||
|
)
|
||||||
|
if use_decode_radix_cache:
|
||||||
# Match prefix against decode's radix cache.
|
# Match prefix against decode's radix cache.
|
||||||
prefix_match = self._match_prefix_and_lock(decode_req.req)
|
prefix_match = self._match_prefix_and_lock(decode_req.req)
|
||||||
prefix_indices = prefix_match.prefix_indices
|
prefix_indices = prefix_match.prefix_indices
|
||||||
@@ -898,7 +971,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
prefix_len = prefix_match.l1_prefix_len
|
prefix_len = prefix_match.l1_prefix_len
|
||||||
total_prefix_len = prefix_match.decode_prefix_len
|
total_prefix_len = prefix_match.decode_prefix_len
|
||||||
|
|
||||||
fill_len = origin_input_len + max(len(decode_req.req.output_ids) - 1, 0)
|
fill_len = self._pre_alloc_fill_len(decode_req.req)
|
||||||
required_alloc_tokens = self._required_alloc_tokens(
|
required_alloc_tokens = self._required_alloc_tokens(
|
||||||
fill_len=fill_len, prefix_len=prefix_len
|
fill_len=fill_len, prefix_len=prefix_len
|
||||||
)
|
)
|
||||||
@@ -915,7 +988,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
prefix_indices = None
|
prefix_indices = None
|
||||||
prefix_len = 0
|
prefix_len = 0
|
||||||
total_prefix_len = 0
|
total_prefix_len = 0
|
||||||
required_alloc_tokens = origin_input_len
|
required_alloc_tokens = self._pre_alloc_fill_len(decode_req.req)
|
||||||
|
|
||||||
required_tokens_for_request = (
|
required_tokens_for_request = (
|
||||||
required_alloc_tokens + self.num_reserved_decode_tokens
|
required_alloc_tokens + self.num_reserved_decode_tokens
|
||||||
@@ -1013,7 +1086,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
.numpy()
|
.numpy()
|
||||||
)
|
)
|
||||||
|
|
||||||
seq_len = len(decode_req.req.origin_input_ids)
|
seq_len = origin_input_len
|
||||||
|
|
||||||
def _mamba_payload():
|
def _mamba_payload():
|
||||||
return [
|
return [
|
||||||
@@ -1108,6 +1181,11 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
state_indices,
|
state_indices,
|
||||||
decode_prefix_len=total_prefix_len,
|
decode_prefix_len=total_prefix_len,
|
||||||
)
|
)
|
||||||
|
if decode_req.is_rebootstrap:
|
||||||
|
self.kv_manager.submit_prefill_recompute(
|
||||||
|
decode_req.kv_receiver,
|
||||||
|
decode_req.req.build_rebootstrap_payload(),
|
||||||
|
)
|
||||||
if (
|
if (
|
||||||
self.transfer_queue.enable_staging
|
self.transfer_queue.enable_staging
|
||||||
and hasattr(decode_req.kv_receiver, "require_staging")
|
and hasattr(decode_req.kv_receiver, "require_staging")
|
||||||
@@ -1341,7 +1419,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
req_pool_indices is not None
|
req_pool_indices is not None
|
||||||
), "req_pool_indices is full! There is a bug in memory estimation."
|
), "req_pool_indices is full! There is a bug in memory estimation."
|
||||||
|
|
||||||
fill_len = len(req.origin_input_ids) + max(len(req.output_ids) - 1, 0)
|
fill_len = self._pre_alloc_fill_len(req)
|
||||||
req.kv_allocated_len = fill_len
|
req.kv_allocated_len = fill_len
|
||||||
req.kv_committed_len = fill_len
|
req.kv_committed_len = fill_len
|
||||||
|
|
||||||
@@ -1584,7 +1662,25 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
|
|||||||
self._commit_hicache_local_restore_to_req(decode_req)
|
self._commit_hicache_local_restore_to_req(decode_req)
|
||||||
|
|
||||||
# Case 3: Success - commit the transfer
|
# Case 3: Success - commit the transfer
|
||||||
decode_req.req.output_ids.append(output_id[0].item())
|
# PD true-retraction rebootstrap: the prefill recomputed the prefix KV
|
||||||
|
# under the current weights and sampled a fresh handoff token, but when
|
||||||
|
# there is a remembered boundary token we are *replaying* an
|
||||||
|
# already-emitted token. Override the handoff with it, and skip
|
||||||
|
# re-committing a logprob for it -- it keeps its original behavior
|
||||||
|
# logprob from before the retract (we never re-score generated tokens
|
||||||
|
# under the new policy). A rebootstrap with no boundary token (retracted
|
||||||
|
# before emitting any output) falls through to the normal path so its
|
||||||
|
# first token and logprob are committed as usual.
|
||||||
|
replayed_boundary = (
|
||||||
|
decode_req.is_rebootstrap
|
||||||
|
and decode_req.req.pd_rebootstrap_forced_output_id is not None
|
||||||
|
)
|
||||||
|
if replayed_boundary:
|
||||||
|
committed_output_id = decode_req.req.pd_rebootstrap_forced_output_id
|
||||||
|
decode_req.req.pd_rebootstrap_forced_output_id = None
|
||||||
|
else:
|
||||||
|
committed_output_id = output_id[0].item()
|
||||||
|
decode_req.req.output_ids.append(committed_output_id)
|
||||||
decode_req.req.cached_tokens = cached_tokens[0].item()
|
decode_req.req.cached_tokens = cached_tokens[0].item()
|
||||||
# The prefill node already reported its prefix-cache hit in
|
# The prefill node already reported its prefix-cache hit in
|
||||||
# cached_tokens[0]. Seed already_computed with it so that
|
# cached_tokens[0]. Seed already_computed with it so that
|
||||||
@@ -1606,7 +1702,7 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
|
|||||||
decode_req.req.output_topk_index = output_topk_index
|
decode_req.req.output_topk_index = output_topk_index
|
||||||
decode_req.req.hidden_states_tensor = output_hidden_states
|
decode_req.req.hidden_states_tensor = output_hidden_states
|
||||||
|
|
||||||
if decode_req.req.return_logprob:
|
if decode_req.req.return_logprob and not replayed_boundary:
|
||||||
decode_req.req.logprob.output_token_logprobs_val.append(
|
decode_req.req.logprob.output_token_logprobs_val.append(
|
||||||
output_token_logprobs_val[0].item()
|
output_token_logprobs_val[0].item()
|
||||||
)
|
)
|
||||||
@@ -1799,9 +1895,9 @@ class SchedulerDisaggregationDecodeMixin:
|
|||||||
# 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)
|
||||||
self.process_decode_queue()
|
|
||||||
if self._engine_paused:
|
if self._engine_paused:
|
||||||
continue
|
continue
|
||||||
|
self.process_decode_queue()
|
||||||
|
|
||||||
# Get the next batch to run
|
# Get the next batch to run
|
||||||
batch = self.get_next_disagg_decode_batch_to_run()
|
batch = self.get_next_disagg_decode_batch_to_run()
|
||||||
@@ -1831,9 +1927,9 @@ class SchedulerDisaggregationDecodeMixin:
|
|||||||
# 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)
|
||||||
self.process_decode_queue()
|
|
||||||
if self._engine_paused:
|
if self._engine_paused:
|
||||||
continue
|
continue
|
||||||
|
self.process_decode_queue()
|
||||||
|
|
||||||
self._apply_war_barrier()
|
self._apply_war_barrier()
|
||||||
|
|
||||||
|
|||||||
@@ -72,6 +72,8 @@ class ScheduleBatchDisaggregationDecodeMixin:
|
|||||||
req.cached_tokens_device += delta
|
req.cached_tokens_device += delta
|
||||||
req.already_computed = seq_len
|
req.already_computed = seq_len
|
||||||
req.is_retracted = False
|
req.is_retracted = False
|
||||||
|
if getattr(req, "pd_rebootstrap_in_progress", False):
|
||||||
|
req.pd_rebootstrap_in_progress = False
|
||||||
pre_lens.append(pre_len)
|
pre_lens.append(pre_len)
|
||||||
|
|
||||||
# Set fields
|
# Set fields
|
||||||
|
|||||||
@@ -474,11 +474,11 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
# 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)
|
||||||
|
if self._engine_paused:
|
||||||
|
continue
|
||||||
self.waiting_queue.extend(
|
self.waiting_queue.extend(
|
||||||
self.disagg_prefill_bootstrap_queue.pop_bootstrapped()
|
self.disagg_prefill_bootstrap_queue.pop_bootstrapped()
|
||||||
)
|
)
|
||||||
if self._engine_paused:
|
|
||||||
continue
|
|
||||||
|
|
||||||
# Get the next batch to run
|
# Get the next batch to run
|
||||||
batch = self.get_next_disagg_prefill_batch_to_run()
|
batch = self.get_next_disagg_prefill_batch_to_run()
|
||||||
@@ -506,11 +506,11 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
# 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)
|
||||||
|
if self._engine_paused:
|
||||||
|
continue
|
||||||
self.waiting_queue.extend(
|
self.waiting_queue.extend(
|
||||||
self.disagg_prefill_bootstrap_queue.pop_bootstrapped()
|
self.disagg_prefill_bootstrap_queue.pop_bootstrapped()
|
||||||
)
|
)
|
||||||
if self._engine_paused:
|
|
||||||
continue
|
|
||||||
|
|
||||||
self._apply_war_barrier()
|
self._apply_war_barrier()
|
||||||
|
|
||||||
|
|||||||
@@ -986,6 +986,10 @@ class Req(ReqDllmMixin):
|
|||||||
self.bootstrap_host: str = bootstrap_host
|
self.bootstrap_host: str = bootstrap_host
|
||||||
self.bootstrap_port: Optional[int] = bootstrap_port
|
self.bootstrap_port: Optional[int] = bootstrap_port
|
||||||
self.bootstrap_room: Optional[int] = bootstrap_room
|
self.bootstrap_room: Optional[int] = bootstrap_room
|
||||||
|
# Decode-local: the already-emitted boundary token to replay when a
|
||||||
|
# retracted request is rebootstrapped. Set in pause_generation(retract)
|
||||||
|
# and consumed in the decode transfer commit; never plumbed to prefill.
|
||||||
|
self.pd_rebootstrap_forced_output_id: Optional[int] = None
|
||||||
self.skip_radix_cache_insert = bootstrap_host == FAKE_BOOTSTRAP_HOST
|
self.skip_radix_cache_insert = bootstrap_host == FAKE_BOOTSTRAP_HOST
|
||||||
self.disagg_kv_sender: Optional[BaseKVSender] = None
|
self.disagg_kv_sender: Optional[BaseKVSender] = None
|
||||||
|
|
||||||
@@ -1503,6 +1507,54 @@ class Req(ReqDllmMixin):
|
|||||||
)
|
)
|
||||||
del self.kv_cache_cpu
|
del self.kv_cache_cpu
|
||||||
|
|
||||||
|
def build_rebootstrap_payload(self) -> dict:
|
||||||
|
"""Build the prefill ``/generate`` payload that asks the original prefill
|
||||||
|
worker to recompute this request's prefix KV under the current weights
|
||||||
|
(PD true-retraction rebootstrap).
|
||||||
|
|
||||||
|
``input_ids`` are coerced to plain ``int`` so the payload is always
|
||||||
|
JSON-serializable even when ``origin_input_ids``/``output_ids`` hold
|
||||||
|
numpy scalars. The sampling-param allow-list forces ``max_new_tokens=1``
|
||||||
|
and drops stop/grammar/min_new_tokens so the recompute only re-derives
|
||||||
|
the prefix KV and samples a single handoff token. The already-emitted
|
||||||
|
boundary token is replayed on the *decode* side (the transfer commit
|
||||||
|
overrides the sampled handoff with it), so it is intentionally not sent
|
||||||
|
to the prefill here.
|
||||||
|
"""
|
||||||
|
# TODO: multi-modal requests are not supported here. The payload only
|
||||||
|
# carries token ``input_ids`` and drops any image/audio/video inputs, so
|
||||||
|
# the rebootstrap recompute would not reproduce the original prefix KV
|
||||||
|
# for multi-modal requests. Add multi-modal support before enabling it.
|
||||||
|
sp = self.sampling_params
|
||||||
|
return {
|
||||||
|
"input_ids": [int(x) for x in self.origin_input_ids]
|
||||||
|
+ [int(x) for x in self.output_ids],
|
||||||
|
"sampling_params": {
|
||||||
|
"max_new_tokens": 1,
|
||||||
|
"temperature": sp.temperature,
|
||||||
|
"top_p": sp.top_p,
|
||||||
|
"top_k": sp.top_k,
|
||||||
|
"min_p": sp.min_p,
|
||||||
|
"frequency_penalty": sp.frequency_penalty,
|
||||||
|
"presence_penalty": sp.presence_penalty,
|
||||||
|
"repetition_penalty": sp.repetition_penalty,
|
||||||
|
"ignore_eos": sp.ignore_eos,
|
||||||
|
"skip_special_tokens": sp.skip_special_tokens,
|
||||||
|
"spaces_between_special_tokens": sp.spaces_between_special_tokens,
|
||||||
|
"no_stop_trim": sp.no_stop_trim,
|
||||||
|
},
|
||||||
|
"return_logprob": False,
|
||||||
|
"stream": False,
|
||||||
|
"rid": self.rid,
|
||||||
|
"bootstrap_host": self.bootstrap_host,
|
||||||
|
"bootstrap_port": self.bootstrap_port,
|
||||||
|
"bootstrap_room": self.bootstrap_room,
|
||||||
|
"priority": self.priority,
|
||||||
|
"extra_key": self.extra_key,
|
||||||
|
"routing_key": self.routing_key,
|
||||||
|
"disagg_prefill_dp_rank": self.disagg_prefill_dp_rank,
|
||||||
|
}
|
||||||
|
|
||||||
def log_time_stats(self):
|
def log_time_stats(self):
|
||||||
# If overlap schedule, we schedule one decode batch ahead so this gets called twice.
|
# If overlap schedule, we schedule one decode batch ahead so this gets called twice.
|
||||||
if self.has_log_time_stats:
|
if self.has_log_time_stats:
|
||||||
@@ -1596,11 +1648,16 @@ def release_req(
|
|||||||
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator,
|
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator,
|
||||||
tree_cache: BasePrefixCache,
|
tree_cache: BasePrefixCache,
|
||||||
hisparse_coordinator: Optional[HiSparseCoordinator],
|
hisparse_coordinator: Optional[HiSparseCoordinator],
|
||||||
|
offload_kv: bool = True,
|
||||||
) -> None:
|
) -> None:
|
||||||
if hisparse_coordinator is not None and not req.finished():
|
if hisparse_coordinator is not None and not req.finished():
|
||||||
hisparse_coordinator.retract_req(req)
|
hisparse_coordinator.retract_req(req)
|
||||||
|
|
||||||
if server_args.disaggregation_mode == "decode":
|
# In decode disaggregation the retracted KV is offloaded to host so it can be
|
||||||
|
# restored later without recompute (see resume_retracted_reqs/load_kv_cache).
|
||||||
|
# Callers that will recompute the KV instead (PD true-retraction rebootstrap)
|
||||||
|
# pass offload_kv=False to skip the wasteful device->host copy.
|
||||||
|
if server_args.disaggregation_mode == "decode" and offload_kv:
|
||||||
req.offload_kv_cache(req_to_token_pool, token_to_kv_pool_allocator)
|
req.offload_kv_cache(req_to_token_pool, token_to_kv_pool_allocator)
|
||||||
# TODO (csy): for preempted requests, we may want to insert into the tree
|
# TODO (csy): for preempted requests, we may want to insert into the tree
|
||||||
release_kv_cache(req, tree_cache, is_insert=False)
|
release_kv_cache(req, tree_cache, is_insert=False)
|
||||||
@@ -1619,6 +1676,7 @@ def retract_all(
|
|||||||
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator,
|
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator,
|
||||||
tree_cache: BasePrefixCache,
|
tree_cache: BasePrefixCache,
|
||||||
hisparse_coordinator: Optional[HiSparseCoordinator],
|
hisparse_coordinator: Optional[HiSparseCoordinator],
|
||||||
|
offload_kv: bool = True,
|
||||||
) -> List[Req]:
|
) -> List[Req]:
|
||||||
retracted_reqs = reqs
|
retracted_reqs = reqs
|
||||||
for idx in range(len(reqs)):
|
for idx in range(len(reqs)):
|
||||||
@@ -1630,6 +1688,7 @@ def retract_all(
|
|||||||
token_to_kv_pool_allocator=token_to_kv_pool_allocator,
|
token_to_kv_pool_allocator=token_to_kv_pool_allocator,
|
||||||
tree_cache=tree_cache,
|
tree_cache=tree_cache,
|
||||||
hisparse_coordinator=hisparse_coordinator,
|
hisparse_coordinator=hisparse_coordinator,
|
||||||
|
offload_kv=offload_kv,
|
||||||
)
|
)
|
||||||
return retracted_reqs
|
return retracted_reqs
|
||||||
|
|
||||||
@@ -2458,7 +2517,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
evict_from_tree_cache(self.tree_cache, num_tokens)
|
evict_from_tree_cache(self.tree_cache, num_tokens)
|
||||||
return self.token_to_kv_pool_allocator.available_size() >= num_tokens
|
return self.token_to_kv_pool_allocator.available_size() >= num_tokens
|
||||||
|
|
||||||
def retract_all(self, server_args: ServerArgs):
|
def retract_all(self, server_args: ServerArgs, offload_kv: bool = True):
|
||||||
retracted_reqs = retract_all(
|
retracted_reqs = retract_all(
|
||||||
reqs=self.reqs,
|
reqs=self.reqs,
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
@@ -2466,6 +2525,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
||||||
tree_cache=self.tree_cache,
|
tree_cache=self.tree_cache,
|
||||||
hisparse_coordinator=self.hisparse_coordinator,
|
hisparse_coordinator=self.hisparse_coordinator,
|
||||||
|
offload_kv=offload_kv,
|
||||||
)
|
)
|
||||||
self.reqs = []
|
self.reqs = []
|
||||||
return retracted_reqs
|
return retracted_reqs
|
||||||
|
|||||||
@@ -4034,9 +4034,22 @@ class Scheduler(
|
|||||||
if recv_req.mode == "retract" and not self.running_batch.is_empty():
|
if recv_req.mode == "retract" and not self.running_batch.is_empty():
|
||||||
self.running_batch.filter_batch()
|
self.running_batch.filter_batch()
|
||||||
if len(self.running_batch.reqs) != 0:
|
if len(self.running_batch.reqs) != 0:
|
||||||
retracted_reqs = self.running_batch.retract_all(self.server_args)
|
# Decode-side retract always rebootstraps (recomputes the KV from
|
||||||
|
# the prefill), so skip the device->host KV offload that release_req
|
||||||
|
# would otherwise do; the offloaded copy would be immediately
|
||||||
|
# discarded. Non-decode modes ignore offload_kv (they never offload).
|
||||||
|
retracted_reqs = self.running_batch.retract_all(
|
||||||
|
self.server_args, offload_kv=False
|
||||||
|
)
|
||||||
for req in retracted_reqs:
|
for req in retracted_reqs:
|
||||||
self._add_request_to_queue(req)
|
if self.disaggregation_mode == DisaggregationMode.DECODE:
|
||||||
|
if req.output_ids:
|
||||||
|
req.pd_rebootstrap_forced_output_id = req.output_ids.pop()
|
||||||
|
req.pd_rebootstrap_in_progress = True
|
||||||
|
req.time_stats.set_retract_time()
|
||||||
|
self.disagg_decode_prealloc_queue.hold_rebootstrap(req)
|
||||||
|
else:
|
||||||
|
self._add_request_to_queue(req)
|
||||||
|
|
||||||
self.running_batch.batch_is_full = False
|
self.running_batch.batch_is_full = False
|
||||||
self.chunked_req = None
|
self.chunked_req = None
|
||||||
@@ -4064,6 +4077,16 @@ class Scheduler(
|
|||||||
f"reserved {before_mb:.1f} MB -> {after_mb:.1f} MB "
|
f"reserved {before_mb:.1f} MB -> {after_mb:.1f} MB "
|
||||||
f"(freed {before_mb - after_mb:.1f} MB)"
|
f"(freed {before_mb - after_mb:.1f} MB)"
|
||||||
)
|
)
|
||||||
|
# Enqueue any rebootstrap requests that were staged during a
|
||||||
|
# retract-mode pause. Deferring until resume keeps the preallocation
|
||||||
|
# queue empty during the pause window (so an intervening weight update
|
||||||
|
# can flush the cache) and recomputes the prefix KV under the updated
|
||||||
|
# weights.
|
||||||
|
if (
|
||||||
|
self.disaggregation_mode == DisaggregationMode.DECODE
|
||||||
|
and self.disagg_decode_prealloc_queue is not None
|
||||||
|
):
|
||||||
|
self.disagg_decode_prealloc_queue.enqueue_held_rebootstrap()
|
||||||
self._engine_paused = False
|
self._engine_paused = False
|
||||||
|
|
||||||
def load_lora_adapter(
|
def load_lora_adapter(
|
||||||
|
|||||||
@@ -1483,11 +1483,11 @@ class ChunkSizePredictor:
|
|||||||
def set_target_latency(self, base_chunk_size: int):
|
def set_target_latency(self, base_chunk_size: int):
|
||||||
"""Set target latency based on base chunk size: target = f(base_chunk_size) - f(0)."""
|
"""Set target latency based on base chunk size: target = f(base_chunk_size) - f(0)."""
|
||||||
|
|
||||||
def f(l: float) -> float:
|
def f(length: float) -> float:
|
||||||
"""Total latency function: f(l) = al^2 + bl + c (or bl + c for linear)"""
|
"""Total latency function: f(length) = a*length^2 + b*length + c."""
|
||||||
return (
|
return (
|
||||||
self.quadratic_coeff_a * l * l
|
self.quadratic_coeff_a * length * length
|
||||||
+ self.linear_coeff_b * l
|
+ self.linear_coeff_b * length
|
||||||
+ self.constant_coeff_c
|
+ self.constant_coeff_c
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -291,6 +291,165 @@ class TestDisaggregationSimulatedRetract(PDDisaggregationServerBase):
|
|||||||
self.assertGreater(metrics["score"], 0.62)
|
self.assertGreater(metrics["score"], 0.62)
|
||||||
|
|
||||||
|
|
||||||
|
class TestDisaggregationPauseResumeDecodeRetract(PDDisaggregationServerBase):
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls):
|
||||||
|
super().setUpClass()
|
||||||
|
cls.model = DEFAULT_MODEL_NAME_FOR_TEST
|
||||||
|
cls.launch_all()
|
||||||
|
|
||||||
|
def test_retract_pause_decode_running_batch(self):
|
||||||
|
"""Retract-mode pause on a disagg decode node must preserve in-flight
|
||||||
|
requests that are already in running_batch."""
|
||||||
|
asyncio.run(self._run_pause_on_decode_running_batch("retract"))
|
||||||
|
|
||||||
|
def test_retract_weight_update_decode_running_batch(self):
|
||||||
|
"""Retract pause + weight update on a disagg decode node.
|
||||||
|
|
||||||
|
This guards the core reason retract mode exists: while paused, the
|
||||||
|
running_batch AND the rebootstrap preallocation queue are empty, so the
|
||||||
|
scheduler is fully idle and the post-update cache flush succeeds (a
|
||||||
|
regression here trips ``assert ..., "Cache flush failed after updating
|
||||||
|
weights"`` and crashes the decode worker). On continue, the retracted
|
||||||
|
requests rebootstrap-recompute their prefix KV under the updated weights
|
||||||
|
and resume to completion.
|
||||||
|
"""
|
||||||
|
asyncio.run(
|
||||||
|
self._run_pause_on_decode_running_batch("retract", weight_update=True)
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _get_decode_num_running_reqs(self, session):
|
||||||
|
"""Query current decode running_batch size from /v1/loads."""
|
||||||
|
async with session.get(
|
||||||
|
self.decode_url + "/v1/loads?include=core",
|
||||||
|
timeout=aiohttp.ClientTimeout(total=5),
|
||||||
|
) as resp:
|
||||||
|
resp.raise_for_status()
|
||||||
|
body = await resp.json()
|
||||||
|
return sum(load["num_running_reqs"] for load in body["loads"])
|
||||||
|
|
||||||
|
async def _wait_for_decode_running_batch(self, session, timeout):
|
||||||
|
deadline = asyncio.get_running_loop().time() + timeout
|
||||||
|
while asyncio.get_running_loop().time() < deadline:
|
||||||
|
if await self._get_decode_num_running_reqs(session) > 0:
|
||||||
|
return
|
||||||
|
await asyncio.sleep(0.2)
|
||||||
|
|
||||||
|
self.fail("Timed out waiting for decode running_batch to become non-empty")
|
||||||
|
|
||||||
|
async def _run_pause_on_decode_running_batch(self, mode, weight_update=False):
|
||||||
|
num_requests = 2
|
||||||
|
max_new_tokens = 512
|
||||||
|
prompt = "Write a detailed numbered explanation of distributed inference. " * 12
|
||||||
|
|
||||||
|
async def _post(session, url, json_data, timeout=30):
|
||||||
|
async with session.post(
|
||||||
|
url,
|
||||||
|
json=json_data,
|
||||||
|
timeout=aiohttp.ClientTimeout(total=timeout),
|
||||||
|
) as resp:
|
||||||
|
resp.raise_for_status()
|
||||||
|
return await resp.json()
|
||||||
|
|
||||||
|
async def _generate(session, request_id):
|
||||||
|
return await _post(
|
||||||
|
session,
|
||||||
|
self.lb_url + "/generate",
|
||||||
|
{
|
||||||
|
"text": f"Request {request_id}: {prompt}",
|
||||||
|
"background": True,
|
||||||
|
"sampling_params": {
|
||||||
|
"temperature": 0,
|
||||||
|
"ignore_eos": True,
|
||||||
|
"max_new_tokens": max_new_tokens,
|
||||||
|
},
|
||||||
|
},
|
||||||
|
timeout=180,
|
||||||
|
)
|
||||||
|
|
||||||
|
async with aiohttp.ClientSession() as session:
|
||||||
|
tasks = [
|
||||||
|
asyncio.create_task(_generate(session, i)) for i in range(num_requests)
|
||||||
|
]
|
||||||
|
decode_paused = False
|
||||||
|
|
||||||
|
try:
|
||||||
|
await self._wait_for_decode_running_batch(session, timeout=30)
|
||||||
|
await asyncio.sleep(0.1)
|
||||||
|
|
||||||
|
self.assertTrue(
|
||||||
|
any(not task.done() for task in tasks),
|
||||||
|
"All requests finished before decode retract pause was issued.",
|
||||||
|
)
|
||||||
|
|
||||||
|
await _post(
|
||||||
|
session,
|
||||||
|
self.decode_url + "/pause_generation",
|
||||||
|
{"mode": mode},
|
||||||
|
)
|
||||||
|
decode_paused = True
|
||||||
|
await asyncio.sleep(1)
|
||||||
|
|
||||||
|
if weight_update:
|
||||||
|
# Reload the same weights from disk while retract-paused. The
|
||||||
|
# update mechanism (disk/tensor/distributed/ipc) is irrelevant
|
||||||
|
# here: they all share flush_cache_after_weight_update, whose
|
||||||
|
# flush asserts the scheduler is fully idle. This must not
|
||||||
|
# crash, proving retracted reqs are not stuck in the prealloc
|
||||||
|
# queue.
|
||||||
|
wu = await _post(
|
||||||
|
session,
|
||||||
|
self.decode_url + "/update_weights_from_disk",
|
||||||
|
{"model_path": self.model},
|
||||||
|
timeout=180,
|
||||||
|
)
|
||||||
|
self.assertTrue(
|
||||||
|
wu.get("success", False),
|
||||||
|
f"update_weights_from_disk failed during retract pause: {wu}",
|
||||||
|
)
|
||||||
|
|
||||||
|
await _post(session, self.decode_url + "/continue_generation", {})
|
||||||
|
decode_paused = False
|
||||||
|
|
||||||
|
responses = await asyncio.wait_for(asyncio.gather(*tasks), timeout=180)
|
||||||
|
finally:
|
||||||
|
if decode_paused:
|
||||||
|
try:
|
||||||
|
await _post(
|
||||||
|
session, self.decode_url + "/continue_generation", {}
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
unfinished = [task for task in tasks if not task.done()]
|
||||||
|
if unfinished:
|
||||||
|
for url in [self.prefill_url, self.decode_url]:
|
||||||
|
try:
|
||||||
|
await _post(
|
||||||
|
session,
|
||||||
|
url + "/abort_request",
|
||||||
|
{"abort_all": True},
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
for task in unfinished:
|
||||||
|
task.cancel()
|
||||||
|
await asyncio.gather(*unfinished, return_exceptions=True)
|
||||||
|
|
||||||
|
for response in responses:
|
||||||
|
self.assertIn("text", response)
|
||||||
|
self.assertGreater(len(response["text"]), 0)
|
||||||
|
|
||||||
|
self.assertGreater(
|
||||||
|
sum(
|
||||||
|
response.get("meta_info", {}).get("num_retractions", 0)
|
||||||
|
for response in responses
|
||||||
|
),
|
||||||
|
0,
|
||||||
|
"Expected pause_generation(retract) to retract a running decode request.",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestDisaggregationPauseResumePrefillLeak(PDDisaggregationServerBase):
|
class TestDisaggregationPauseResumePrefillLeak(PDDisaggregationServerBase):
|
||||||
"""Regression test: pause_generation must not leak prefill requests into
|
"""Regression test: pause_generation must not leak prefill requests into
|
||||||
running_batch. With a small --max-running-requests the leak fills the
|
running_batch. With a small --max-running-requests the leak fills the
|
||||||
|
|||||||
@@ -165,9 +165,13 @@ class TestRegisterToBootstrap(CustomTestCase):
|
|||||||
"rank_port",
|
"rank_port",
|
||||||
"page_size",
|
"page_size",
|
||||||
"kv_cache_dtype",
|
"kv_cache_dtype",
|
||||||
|
# Self-registered HTTP API port used to derive the PD retract
|
||||||
|
# rebootstrap /generate URL on the decode side.
|
||||||
|
"prefill_http_port",
|
||||||
]
|
]
|
||||||
for field in required_fields:
|
for field in required_fields:
|
||||||
self.assertIn(field, payload)
|
self.assertIn(field, payload)
|
||||||
|
self.assertEqual(payload["prefill_http_port"], 30000)
|
||||||
|
|
||||||
@patch("sglang.srt.disaggregation.common.conn.time")
|
@patch("sglang.srt.disaggregation.common.conn.time")
|
||||||
@patch("sglang.srt.disaggregation.common.conn.requests.put")
|
@patch("sglang.srt.disaggregation.common.conn.requests.put")
|
||||||
@@ -266,6 +270,7 @@ class TestRegisterToBootstrap(CustomTestCase):
|
|||||||
mgr.server_args = MagicMock()
|
mgr.server_args = MagicMock()
|
||||||
mgr.server_args.kv_cache_dtype = "auto"
|
mgr.server_args.kv_cache_dtype = "auto"
|
||||||
mgr.server_args.load_balance_method = "follow_bootstrap_room"
|
mgr.server_args.load_balance_method = "follow_bootstrap_room"
|
||||||
|
mgr.server_args.port = 30000
|
||||||
|
|
||||||
return mgr
|
return mgr
|
||||||
|
|
||||||
|
|||||||
@@ -1,8 +1,11 @@
|
|||||||
|
import json
|
||||||
import sys
|
import sys
|
||||||
|
import threading
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
from unittest.mock import MagicMock, patch
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.disaggregation.decode import ( # noqa: E402
|
from sglang.srt.disaggregation.decode import ( # noqa: E402
|
||||||
@@ -10,7 +13,7 @@ from sglang.srt.disaggregation.decode import ( # noqa: E402
|
|||||||
SchedulerDisaggregationDecodeMixin,
|
SchedulerDisaggregationDecodeMixin,
|
||||||
)
|
)
|
||||||
from sglang.srt.disaggregation.utils import DisaggregationMode # noqa: E402
|
from sglang.srt.disaggregation.utils import DisaggregationMode # noqa: E402
|
||||||
from sglang.srt.managers.schedule_batch import FINISH_ABORT # noqa: E402
|
from sglang.srt.managers.schedule_batch import FINISH_ABORT, Req # noqa: E402
|
||||||
from sglang.srt.managers.scheduler import Scheduler # noqa: E402
|
from sglang.srt.managers.scheduler import Scheduler # noqa: E402
|
||||||
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
||||||
|
|
||||||
@@ -96,6 +99,7 @@ class TestDecodePreallocQueuePriority(unittest.TestCase):
|
|||||||
waiting_for_input=True,
|
waiting_for_input=True,
|
||||||
kv_receiver=MagicMock(),
|
kv_receiver=MagicMock(),
|
||||||
metadata_buffer_index=-1,
|
metadata_buffer_index=-1,
|
||||||
|
is_rebootstrap=False,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _new_queue(self, decode_reqs, *, low_priority_values_first: bool = False):
|
def _new_queue(self, decode_reqs, *, low_priority_values_first: bool = False):
|
||||||
@@ -205,6 +209,204 @@ class TestDecodePreallocQueuePriority(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestDecodePreallocQueueRebootstrapPayload(unittest.TestCase):
|
||||||
|
"""The decode scheduler builds the rebootstrap ``/generate`` payload; the
|
||||||
|
dispatch itself now lives on the kv manager (see
|
||||||
|
``TestCommonKVManagerPrefillRecompute``)."""
|
||||||
|
|
||||||
|
def _sampling_params(self):
|
||||||
|
return SimpleNamespace(
|
||||||
|
temperature=0.0,
|
||||||
|
top_p=1.0,
|
||||||
|
top_k=-1,
|
||||||
|
min_p=0.0,
|
||||||
|
frequency_penalty=0.0,
|
||||||
|
presence_penalty=0.0,
|
||||||
|
repetition_penalty=1.0,
|
||||||
|
ignore_eos=False,
|
||||||
|
skip_special_tokens=True,
|
||||||
|
spaces_between_special_tokens=True,
|
||||||
|
no_stop_trim=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _new_req(self):
|
||||||
|
return SimpleNamespace(
|
||||||
|
rid="rid-0",
|
||||||
|
origin_input_ids=np.array([1, 2], dtype=np.int32),
|
||||||
|
output_ids=[np.int32(3), np.int32(4)],
|
||||||
|
sampling_params=self._sampling_params(),
|
||||||
|
bootstrap_host="127.0.0.1",
|
||||||
|
bootstrap_port=30000,
|
||||||
|
bootstrap_room=7,
|
||||||
|
priority=10,
|
||||||
|
extra_key=None,
|
||||||
|
routing_key=None,
|
||||||
|
disagg_prefill_dp_rank=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_build_rebootstrap_payload_converts_numpy_ids_to_json_lists(self):
|
||||||
|
req = self._new_req()
|
||||||
|
|
||||||
|
# build_rebootstrap_payload lives on Req; exercise it unbound with a
|
||||||
|
# namespace that carries the attributes it reads.
|
||||||
|
payload = Req.build_rebootstrap_payload(req)
|
||||||
|
|
||||||
|
# origin_input_ids + output_ids, coerced to plain python ints.
|
||||||
|
self.assertEqual(payload["input_ids"], [1, 2, 3, 4])
|
||||||
|
self.assertTrue(all(type(x) is int for x in payload["input_ids"]))
|
||||||
|
self.assertEqual(payload["sampling_params"]["max_new_tokens"], 1)
|
||||||
|
self.assertEqual(payload["bootstrap_room"], 7)
|
||||||
|
# The prefill /generate URL is derived from bootstrap info on the decode
|
||||||
|
# side, not sent in the payload; and the boundary token is replayed via
|
||||||
|
# the decode-side override, so neither belongs in the payload.
|
||||||
|
self.assertNotIn("pd_rebootstrap_prefill_url", payload)
|
||||||
|
self.assertNotIn("pd_rebootstrap_forced_output_id", payload)
|
||||||
|
# Must be JSON-serializable (numpy scalars would raise here).
|
||||||
|
json.dumps(payload)
|
||||||
|
|
||||||
|
|
||||||
|
class TestCommonKVManagerPrefillRecompute(unittest.TestCase):
|
||||||
|
"""The kv manager owns the shared executor + HTTP session and routes any
|
||||||
|
rebootstrap ``/generate`` failure through ``kv_receiver.abort()`` ->
|
||||||
|
``KVPoll.Failed`` so the scheduler's normal transfer-failure streaming runs.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def _new_manager(self):
|
||||||
|
from sglang.srt.disaggregation.common.conn import CommonKVManager
|
||||||
|
|
||||||
|
mgr = CommonKVManager.__new__(CommonKVManager)
|
||||||
|
mgr._prefill_recompute_executor = None
|
||||||
|
mgr._prefill_recompute_executor_lock = threading.Lock()
|
||||||
|
mgr._prefill_recompute_sessions = threading.local()
|
||||||
|
mgr.waiting_timeout = 300
|
||||||
|
mgr.failure_records = {}
|
||||||
|
mgr.failure_lock = threading.Lock()
|
||||||
|
# Only the attn-tp/attn-cp group leader on the first PP stage issues the
|
||||||
|
# single rebootstrap /generate; default the mock manager to that leader.
|
||||||
|
mgr.attn_tp_rank = 0
|
||||||
|
mgr.attn_cp_rank = 0
|
||||||
|
mgr.pp_rank = 0
|
||||||
|
# Decode-side prefill info cache; the rebootstrap /generate URL is derived
|
||||||
|
# from here (bootstrap_addr host + self-registered prefill_http_port)
|
||||||
|
# instead of a router-injected pd_rebootstrap_prefill_url.
|
||||||
|
mgr.prefill_info_table = {}
|
||||||
|
return mgr
|
||||||
|
|
||||||
|
def _register_prefill_info(self, mgr, bootstrap_addr, http_port):
|
||||||
|
from sglang.srt.disaggregation.common.conn import PrefillServerInfo
|
||||||
|
|
||||||
|
mgr.prefill_info_table[bootstrap_addr] = PrefillServerInfo(
|
||||||
|
attn_tp_size=1,
|
||||||
|
attn_cp_size=1,
|
||||||
|
dp_size=1,
|
||||||
|
pp_size=1,
|
||||||
|
page_size=1,
|
||||||
|
kv_cache_dtype=None,
|
||||||
|
follow_bootstrap_room=True,
|
||||||
|
prefill_http_port=http_port,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _payload(self):
|
||||||
|
return {
|
||||||
|
"input_ids": [1, 2, 3, 4],
|
||||||
|
"rid": "rid-0",
|
||||||
|
}
|
||||||
|
|
||||||
|
def test_submit_dispatches_run_to_shared_executor(self):
|
||||||
|
mgr = self._new_manager()
|
||||||
|
mgr._prefill_recompute_executor = MagicMock()
|
||||||
|
receiver = MagicMock(bootstrap_room=7, bootstrap_addr="127.0.0.1:8998")
|
||||||
|
self._register_prefill_info(mgr, "127.0.0.1:8998", 30000)
|
||||||
|
|
||||||
|
mgr.submit_prefill_recompute(receiver, self._payload())
|
||||||
|
|
||||||
|
mgr._prefill_recompute_executor.submit.assert_called_once()
|
||||||
|
args = mgr._prefill_recompute_executor.submit.call_args[0]
|
||||||
|
self.assertEqual(args[0], mgr._run_prefill_recompute)
|
||||||
|
self.assertIs(args[1], receiver)
|
||||||
|
# URL derived from bootstrap_addr host + registered prefill_http_port.
|
||||||
|
self.assertEqual(args[2], "http://127.0.0.1:30000")
|
||||||
|
receiver.abort.assert_not_called()
|
||||||
|
|
||||||
|
def test_submit_is_noop_on_non_leader_ranks(self):
|
||||||
|
# A retracted request is replicated across every rank in its attention
|
||||||
|
# TP/CP group and every PP stage; only the group/first-stage leader must
|
||||||
|
# POST the single /generate, or the prefill recomputes it once per rank.
|
||||||
|
for attn_tp_rank, attn_cp_rank, pp_rank in (
|
||||||
|
(1, 0, 0),
|
||||||
|
(0, 1, 0),
|
||||||
|
(0, 0, 1),
|
||||||
|
):
|
||||||
|
with self.subTest(
|
||||||
|
attn_tp_rank=attn_tp_rank,
|
||||||
|
attn_cp_rank=attn_cp_rank,
|
||||||
|
pp_rank=pp_rank,
|
||||||
|
):
|
||||||
|
mgr = self._new_manager()
|
||||||
|
mgr.attn_tp_rank = attn_tp_rank
|
||||||
|
mgr.attn_cp_rank = attn_cp_rank
|
||||||
|
mgr.pp_rank = pp_rank
|
||||||
|
mgr._prefill_recompute_executor = MagicMock()
|
||||||
|
receiver = MagicMock(bootstrap_room=7, bootstrap_addr="127.0.0.1:8998")
|
||||||
|
self._register_prefill_info(mgr, "127.0.0.1:8998", 30000)
|
||||||
|
|
||||||
|
mgr.submit_prefill_recompute(receiver, self._payload())
|
||||||
|
|
||||||
|
mgr._prefill_recompute_executor.submit.assert_not_called()
|
||||||
|
receiver.abort.assert_not_called()
|
||||||
|
self.assertEqual(mgr.failure_records, {})
|
||||||
|
|
||||||
|
def test_submit_unresolved_url_fails_via_abort(self):
|
||||||
|
mgr = self._new_manager()
|
||||||
|
mgr._prefill_recompute_executor = MagicMock()
|
||||||
|
# No prefill_info registered for this bootstrap_addr -> URL unresolved.
|
||||||
|
receiver = MagicMock(bootstrap_room=7, bootstrap_addr="127.0.0.1:8998")
|
||||||
|
|
||||||
|
mgr.submit_prefill_recompute(receiver, self._payload())
|
||||||
|
|
||||||
|
receiver.abort.assert_called_once()
|
||||||
|
mgr._prefill_recompute_executor.submit.assert_not_called()
|
||||||
|
self.assertIn(7, mgr.failure_records)
|
||||||
|
|
||||||
|
def test_run_aborts_on_http_error(self):
|
||||||
|
mgr = self._new_manager()
|
||||||
|
session = MagicMock()
|
||||||
|
session.post.return_value = SimpleNamespace(status_code=500, text="boom")
|
||||||
|
mgr._prefill_recompute_sessions.session = session
|
||||||
|
receiver = MagicMock(bootstrap_room=7)
|
||||||
|
|
||||||
|
mgr._run_prefill_recompute(receiver, "http://prefill", self._payload())
|
||||||
|
|
||||||
|
session.post.assert_called_once()
|
||||||
|
receiver.abort.assert_called_once()
|
||||||
|
self.assertIn(7, mgr.failure_records)
|
||||||
|
|
||||||
|
def test_run_aborts_on_exception(self):
|
||||||
|
mgr = self._new_manager()
|
||||||
|
session = MagicMock()
|
||||||
|
session.post.side_effect = RuntimeError("network down")
|
||||||
|
mgr._prefill_recompute_sessions.session = session
|
||||||
|
receiver = MagicMock(bootstrap_room=7)
|
||||||
|
|
||||||
|
mgr._run_prefill_recompute(receiver, "http://prefill", self._payload())
|
||||||
|
|
||||||
|
receiver.abort.assert_called_once()
|
||||||
|
self.assertIn(7, mgr.failure_records)
|
||||||
|
|
||||||
|
def test_run_success_does_not_abort(self):
|
||||||
|
mgr = self._new_manager()
|
||||||
|
session = MagicMock()
|
||||||
|
session.post.return_value = SimpleNamespace(status_code=200, text="")
|
||||||
|
mgr._prefill_recompute_sessions.session = session
|
||||||
|
receiver = MagicMock(bootstrap_room=7)
|
||||||
|
|
||||||
|
mgr._run_prefill_recompute(receiver, "http://prefill", self._payload())
|
||||||
|
|
||||||
|
session.post.assert_called_once()
|
||||||
|
receiver.abort.assert_not_called()
|
||||||
|
self.assertEqual(mgr.failure_records, {})
|
||||||
|
|
||||||
|
|
||||||
class TestDecodePrebuiltPriority(unittest.TestCase):
|
class TestDecodePrebuiltPriority(unittest.TestCase):
|
||||||
def test_waiting_queue_is_sorted_before_prebuilt_selection(self):
|
def test_waiting_queue_is_sorted_before_prebuilt_selection(self):
|
||||||
scheduler = Scheduler.__new__(Scheduler)
|
scheduler = Scheduler.__new__(Scheduler)
|
||||||
@@ -229,8 +431,8 @@ class TestDecodePrebuiltPriority(unittest.TestCase):
|
|||||||
)
|
)
|
||||||
scheduler.future_map = MagicMock()
|
scheduler.future_map = MagicMock()
|
||||||
scheduler.policy = MagicMock()
|
scheduler.policy = MagicMock()
|
||||||
scheduler.policy.calc_priority.side_effect = (
|
scheduler.policy.calc_priority.side_effect = lambda waiting_queue, _: (
|
||||||
lambda waiting_queue, _: waiting_queue.sort(key=lambda req: -req.priority)
|
waiting_queue.sort(key=lambda req: -req.priority)
|
||||||
)
|
)
|
||||||
|
|
||||||
new_batch = MagicMock()
|
new_batch = MagicMock()
|
||||||
|
|||||||
@@ -1,5 +1,6 @@
|
|||||||
import unittest
|
import unittest
|
||||||
from collections import deque
|
from collections import deque
|
||||||
|
from types import SimpleNamespace
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
@@ -7,7 +8,11 @@ from sglang.test.test_utils import maybe_stub_sgl_kernel
|
|||||||
|
|
||||||
maybe_stub_sgl_kernel()
|
maybe_stub_sgl_kernel()
|
||||||
|
|
||||||
from sglang.srt.managers.io_struct import PauseGenerationReqInput
|
from sglang.srt.disaggregation.utils import DisaggregationMode
|
||||||
|
from sglang.srt.managers.io_struct import (
|
||||||
|
ContinueGenerationReqInput,
|
||||||
|
PauseGenerationReqInput,
|
||||||
|
)
|
||||||
from sglang.srt.managers.scheduler import Scheduler
|
from sglang.srt.managers.scheduler import Scheduler
|
||||||
from sglang.srt.managers.scheduler_components.pool_stats_observer import PoolStats
|
from sglang.srt.managers.scheduler_components.pool_stats_observer import PoolStats
|
||||||
|
|
||||||
@@ -31,6 +36,7 @@ class TestSchedulerPauseGeneration(unittest.TestCase):
|
|||||||
scheduler.tree_cache.protected_size.return_value = 0
|
scheduler.tree_cache.protected_size.return_value = 0
|
||||||
scheduler.req_to_token_pool = MagicMock()
|
scheduler.req_to_token_pool = MagicMock()
|
||||||
scheduler.result_queue = deque()
|
scheduler.result_queue = deque()
|
||||||
|
scheduler.disaggregation_mode = DisaggregationMode.NULL
|
||||||
# Support _kv_snap diagnostic logging in patched schedulers
|
# Support _kv_snap diagnostic logging in patched schedulers
|
||||||
scheduler.token_to_kv_pool_allocator = MagicMock()
|
scheduler.token_to_kv_pool_allocator = MagicMock()
|
||||||
scheduler.token_to_kv_pool_allocator.available_size.return_value = 1000
|
scheduler.token_to_kv_pool_allocator.available_size.return_value = 1000
|
||||||
@@ -127,6 +133,53 @@ class TestSchedulerPauseGeneration(unittest.TestCase):
|
|||||||
self.assertEqual(scheduler._add_request_to_queue.call_count, 2)
|
self.assertEqual(scheduler._add_request_to_queue.call_count, 2)
|
||||||
self.assertIsNone(scheduler.chunked_req)
|
self.assertIsNone(scheduler.chunked_req)
|
||||||
|
|
||||||
|
def test_pd_decode_retract_requeues_for_rebootstrap(self):
|
||||||
|
"""PD decode retract should rebootstrap instead of resuming stale CPU KV."""
|
||||||
|
scheduler = self._new_scheduler()
|
||||||
|
scheduler.disaggregation_mode = DisaggregationMode.DECODE
|
||||||
|
scheduler.last_batch = None
|
||||||
|
scheduler.running_batch.reqs = [MagicMock()]
|
||||||
|
scheduler.running_batch.is_empty.return_value = False
|
||||||
|
scheduler._add_request_to_queue = MagicMock()
|
||||||
|
scheduler.disagg_decode_prealloc_queue = MagicMock()
|
||||||
|
|
||||||
|
req = SimpleNamespace(
|
||||||
|
output_ids=[10, 11, 12],
|
||||||
|
time_stats=MagicMock(),
|
||||||
|
)
|
||||||
|
scheduler.running_batch.retract_all.return_value = [req]
|
||||||
|
scheduler.running_batch.filter_batch = MagicMock()
|
||||||
|
scheduler.server_args = MagicMock()
|
||||||
|
|
||||||
|
scheduler.pause_generation(PauseGenerationReqInput(mode="retract"))
|
||||||
|
|
||||||
|
scheduler._add_request_to_queue.assert_not_called()
|
||||||
|
scheduler.disagg_decode_prealloc_queue.hold_rebootstrap.assert_called_once_with(
|
||||||
|
req
|
||||||
|
)
|
||||||
|
self.assertEqual(req.output_ids, [10, 11])
|
||||||
|
self.assertEqual(req.pd_rebootstrap_forced_output_id, 12)
|
||||||
|
self.assertTrue(req.pd_rebootstrap_in_progress)
|
||||||
|
# Rebootstrap recomputes the KV from the prefill, so the retract must skip
|
||||||
|
# the device->host KV offload rather than offload-then-delete it.
|
||||||
|
scheduler.running_batch.retract_all.assert_called_once_with(
|
||||||
|
scheduler.server_args, offload_kv=False
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_pd_decode_continue_releases_held_rebootstrap(self):
|
||||||
|
"""continue_generation must enqueue staged rebootstrap reqs on resume."""
|
||||||
|
scheduler = self._new_scheduler()
|
||||||
|
scheduler.disaggregation_mode = DisaggregationMode.DECODE
|
||||||
|
scheduler.disagg_decode_prealloc_queue = MagicMock()
|
||||||
|
scheduler._engine_paused = True
|
||||||
|
|
||||||
|
scheduler.continue_generation(
|
||||||
|
ContinueGenerationReqInput(torch_empty_cache=False)
|
||||||
|
)
|
||||||
|
|
||||||
|
scheduler.disagg_decode_prealloc_queue.enqueue_held_rebootstrap.assert_called_once_with()
|
||||||
|
self.assertFalse(scheduler._engine_paused)
|
||||||
|
|
||||||
def test_abort_drains_overlap_queue(self):
|
def test_abort_drains_overlap_queue(self):
|
||||||
"""abort with overlap enabled should drain the result_queue."""
|
"""abort with overlap enabled should drain the result_queue."""
|
||||||
scheduler = self._new_scheduler()
|
scheduler = self._new_scheduler()
|
||||||
|
|||||||
@@ -310,6 +310,10 @@ class TestDecodeLockRefScenarios(unittest.TestCase):
|
|||||||
decode_req = MagicMock()
|
decode_req = MagicMock()
|
||||||
decode_req.req = req
|
decode_req.req = req
|
||||||
decode_req.waiting_for_input = True
|
decode_req.waiting_for_input = True
|
||||||
|
# Non-rebootstrap request: exercise the normal decode radix-cache path
|
||||||
|
# (a truthy MagicMock would disable use_decode_radix_cache via the
|
||||||
|
# `not decode_req.is_rebootstrap` gate in pop_preallocated).
|
||||||
|
decode_req.is_rebootstrap = False
|
||||||
|
|
||||||
queue.queue = [decode_req]
|
queue.queue = [decode_req]
|
||||||
queue.pending_reqs = []
|
queue.pending_reqs = []
|
||||||
|
|||||||
Reference in New Issue
Block a user