Complete server warmup before scripted runtime scripts start (#27445)

This commit is contained in:
fzyzcjy
2026-06-07 17:37:37 +08:00
committed by GitHub
parent a39c428d3f
commit 14b8f98a21
3 changed files with 70 additions and 6 deletions
@@ -151,12 +151,13 @@ class ScriptedHttpServer:
# _reset_engine_state, which POSTs to this server's own HTTP port, so
# block until the port is bound and routing before any script can run.
#
# Wait for *any* HTTP response, not status 200: in scripted mode the
# scheduler is driven by the script, so normal warmup never completes
# and /health stays 503 (server_status == Starting) for the whole run.
# A 503 still proves the socket is bound and routes are registered,
# which is all the control POSTs need.
url = f"{self._base_url}/health"
# Poll /model_info rather than /health: once the server reports Up,
# /health runs the generation-based health check (a real probe request
# through the scheduler), which cannot complete while the scheduler
# waits for scripts between RunScript commands. Any /model_info
# response proves the socket is bound and routes are registered, which
# is all the control POSTs need.
url = f"{self._base_url}/model_info"
deadline = time.monotonic() + HTTP_READY_TIMEOUT_S
while time.monotonic() < deadline:
if not self._server_process.is_alive():
@@ -2,6 +2,7 @@ from __future__ import annotations
import logging
import sys
import time
import traceback
from dataclasses import dataclass
from pathlib import Path
@@ -36,6 +37,9 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__)
RESET_DRAIN_MAX_STEPS: int = 200
# Below the test-side LISTENER_ACCEPT_TIMEOUT_S so a stuck warmup surfaces as
# this specific error instead of a generic handshake timeout.
WARMUP_DRIVE_TIMEOUT_S: float = 120.0
@dataclass(frozen=True, slots=True)
@@ -47,6 +51,47 @@ class ScriptedBatchRecord:
chunked_rid: Optional[str]
def _drive_engine_through_warmup(ctx: ScriptedContext) -> Generator:
"""Run the engine until the server warmup request has been received and
fully processed, so scripts never observe foreign warmup traffic."""
scheduler = ctx.scheduler
server_args = scheduler.server_args
if server_args.skip_server_warmup:
logger.info("scripted_runtime: skip_server_warmup set, not driving warmup")
return
logger.info("scripted_runtime: driving engine until server warmup completes")
start_time = time.monotonic()
# is_fully_idle() can transiently report idle while a PP microbatch result
# is still in flight, so require it to hold for two full microbatch
# rotations after the warmup request was observed on the recv socket.
quiesce_iters = 2 * (server_args.pp_size + server_args.pp_async_batch_depth)
proxy = ctx._tokenizer_recv_proxy
deadline = start_time + WARMUP_DRIVE_TIMEOUT_S
idle_streak = 0
while idle_streak < quiesce_iters:
if time.monotonic() >= deadline:
raise RuntimeError(
"scripted_runtime: server warmup did not complete within "
f"{WARMUP_DRIVE_TIMEOUT_S}s "
f"(work_reqs_seen={proxy.work_reqs_seen}, "
f"idle_streak={idle_streak})"
)
yield
if proxy.work_reqs_seen > 0 and scheduler.is_fully_idle():
idle_streak += 1
else:
idle_streak = 0
logger.info(
"scripted_runtime: server warmup drained in %.1fs (work_reqs_seen=%d)",
time.monotonic() - start_time,
proxy.work_reqs_seen,
)
def _reset_engine_state(ctx: ScriptedContext) -> Generator:
scheduler = ctx.scheduler
@@ -106,6 +151,7 @@ class ScriptedSchedulerHook:
ctx_zmq = zmq.Context()
socket = get_zmq_socket(ctx_zmq, zmq.PAIR, endpoint, bind=False)
try:
yield from _drive_engine_through_warmup(self._context)
socket.send_pyobj(HookReady())
while True:
msg = socket.recv_pyobj()
@@ -6,12 +6,27 @@ from typing import Any, Callable
import zmq
from sglang.srt.managers.io_struct import (
BatchTokenizedEmbeddingReqInput,
BatchTokenizedGenerateReqInput,
TokenizedEmbeddingReqInput,
TokenizedGenerateReqInput,
)
_WORK_REQ_TYPES = (
TokenizedGenerateReqInput,
TokenizedEmbeddingReqInput,
BatchTokenizedGenerateReqInput,
BatchTokenizedEmbeddingReqInput,
)
class ScriptedTokenizerRecvProxy:
def __init__(self, *, underlying: zmq.Socket) -> None:
self._underlying = underlying
self._buffer: deque = deque()
self.work_reqs_seen: int = 0
def recv_pyobj(self, flags: int = 0) -> Any:
self._drain_underlying()
@@ -52,4 +67,6 @@ class ScriptedTokenizerRecvProxy:
req = self._underlying.recv_pyobj(zmq.NOBLOCK)
except zmq.ZMQError:
break
if isinstance(req, _WORK_REQ_TYPES):
self.work_reqs_seen += 1
self._buffer.append(req)