[diffusion] feat: add --warmup-mode enum server arg (#28184)

This commit is contained in:
Mick
2026-06-14 23:09:04 +08:00
committed by GitHub
parent 582bd23f71
commit ec36dde580
4 changed files with 197 additions and 11 deletions
@@ -29,8 +29,9 @@ def add_multimodal_gen_serve_args(parser: argparse.ArgumentParser):
def execute_serve_cmd(args: argparse.Namespace, unknown_args: list[str] | None = None):
"""The entry point for the serve command."""
# use server-based warmup for production
server_args = ServerArgs.from_cli_args(
args, unknown_args, default_args={"warmup": True, "server_warmup": True}
args, unknown_args, default_args={"warmup_mode": "server"}
)
dispatch_launch(server_args)
@@ -112,6 +112,9 @@ class Backend(str, Enum):
return [backend.value for backend in cls]
WARMUP_MODES = ("off", "request", "server")
@dataclasses.dataclass
class ServerArgs(DisaggServerArgsMixin):
# Model and path configuration (for convenience)
@@ -223,8 +226,20 @@ class ServerArgs(DisaggServerArgsMixin):
enable_layerwise_nvtx_marker: bool = False
# warmup
# `warmup_mode` is the canonical knob: one of WARMUP_MODES
# - "off": no warmup.
# - "server": server-based warmup — a synthetic request right after the
# server is ready, before real traffic
# - "request": request-based warmup — warm on the first real request(s).
# This is a BENCHMARK aid
# existing consumers keep working) and as deprecated CLI aliases. None means
# "derive the mode from the legacy booleans"; _adjust_warmup resolves it.
warmup_mode: str | None = None
# deprecated: warmup and server_warmup
warmup: bool = False
server_warmup: bool = False
warmup_resolutions: list[str] = None
warmup_steps: int = 1
@@ -689,6 +704,28 @@ class ServerArgs(DisaggServerArgsMixin):
return None, None
def _adjust_warmup(self):
# --warmup-mode > --warmup/--server-warmup
mode_explicit = self.is_arg_explicitly_set("warmup_mode")
legacy_explicit = self.is_arg_explicitly_set(
"warmup"
) or self.is_arg_explicitly_set("server_warmup")
if self.warmup_mode is not None:
if self.warmup_mode not in WARMUP_MODES:
raise ValueError(
f"Invalid --warmup-mode {self.warmup_mode!r}; "
f"expected one of {WARMUP_MODES}."
)
if mode_explicit and legacy_explicit:
logger.warning(
"Both --warmup-mode and the deprecated --warmup/--server-warmup "
"were set; --warmup-mode=%s takes precedence.",
self.warmup_mode,
)
if mode_explicit or not legacy_explicit:
self.warmup = self.warmup_mode != "off"
self.server_warmup = self.warmup_mode == "server"
# Explicit resolutions imply warmup is on (request-based).
if self.warmup_resolutions is not None:
self.warmup = True
@@ -698,6 +735,10 @@ class ServerArgs(DisaggServerArgsMixin):
if not self.warmup:
self.server_warmup = False
self.warmup_mode = (
"off" if not self.warmup else "server" if self.server_warmup else "request"
)
@staticmethod
def _require_port(port: int, name: str) -> None:
"""Raise if *port* is occupied (used under ``--strict-ports``)."""
@@ -1251,19 +1292,32 @@ class ServerArgs(DisaggServerArgsMixin):
)
# warmup
parser.add_argument(
"--warmup-mode",
type=str,
choices=list(WARMUP_MODES),
default=ServerArgs.warmup_mode,
help=(
"Warmup mode (canonical knob). One of: "
"`off` (no warmup); `request` (request-based: warm on real "
"incoming requests); `server` (server-based: a synthetic warmup "
"request right after the server is ready, before traffic). "
"Takes precedence over the deprecated --warmup/--server-warmup. "
"`sglang serve` defaults to `server`; other entrypoints default "
"to request-based when warmup is enabled. When enabled, look for "
"the line ending with `(with warmup excluded)` for actual "
"processing time."
),
)
parser.add_argument(
"--warmup",
action=StoreBoolean,
default=ServerArgs.warmup,
help=(
"Perform warmup before normal traffic. `sglang serve` runs a "
"server warmup through the scheduler client after HTTP is ready. "
"Other client entrypoints run explicit `--warmup-resolutions` "
"through the scheduler client, otherwise they use request-based "
"warmup. Recommended to enable when benchmarking to ensure fair "
"comparison and best performance. When enabled, look for the "
"line ending with `(with warmup excluded)` for actual processing "
"time."
"[DEPRECATED: use --warmup-mode] Perform warmup before normal "
"traffic. Maps to --warmup-mode request (or server, combined "
"with --server-warmup). Recommended when benchmarking for fair "
"comparison and best performance."
),
)
parser.add_argument(
@@ -1283,7 +1337,10 @@ class ServerArgs(DisaggServerArgsMixin):
"--server-warmup",
action=StoreBoolean,
default=ServerArgs.server_warmup,
help="Send a warmup request after server ready",
help=(
"[DEPRECATED: use --warmup-mode server] Send a synthetic warmup "
"request after the server is ready (server-based warmup)."
),
)
# layerwise offload
@@ -21,7 +21,9 @@ from sglang.multimodal_gen.runtime.warmup_request_builder import (
logger = init_logger(__name__)
MINIMUM_PICTURE_BASE64_FOR_WARMUP = "data:image/jpg;base64,iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg=="
# a 64x64 image because some pipelines reject smaller inputs (e.g. FLUX.2's
# diffusers image processor requires both dimensions >= 64px)
MINIMUM_PICTURE_BASE64_FOR_WARMUP = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAEAAAABACAIAAAAlC+aJAAAAS0lEQVR42u3PMQ0AAAwDoEqv9ErYvQQckD4XAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAYHLAB8+AWnmfUycAAAAAElFTkSuQmCC"
def get_first_generation_req(req_or_group: Any) -> Req | None:
@@ -496,6 +496,132 @@ class TestServerArgsPathExpansion(unittest.TestCase):
self.assertFalse(server_args.server_warmup)
class TestWarmupModeNormalization(unittest.TestCase):
"""`_adjust_warmup` resolves the canonical warmup_mode and its derived booleans."""
def _resolve(
self,
*,
warmup_mode=None,
warmup=False,
server_warmup=False,
warmup_resolutions=None,
disagg_role=None,
explicit=(),
):
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
sa = ServerArgs.__new__(ServerArgs)
sa.warmup_mode = warmup_mode
sa.warmup = warmup
sa.server_warmup = server_warmup
sa.warmup_resolutions = warmup_resolutions
sa.disagg_role = RoleType.MONOLITHIC if disagg_role is None else disagg_role
sa._explicit_arg_names = set(explicit)
sa._adjust_warmup()
return sa
def test_explicit_mode_off_disables_all(self):
sa = self._resolve(warmup_mode="off", explicit=("warmup_mode",))
self.assertEqual(sa.warmup_mode, "off")
self.assertFalse(sa.warmup)
self.assertFalse(sa.server_warmup)
def test_explicit_mode_request(self):
sa = self._resolve(warmup_mode="request", explicit=("warmup_mode",))
self.assertEqual(sa.warmup_mode, "request")
self.assertTrue(sa.warmup)
self.assertFalse(sa.server_warmup)
def test_explicit_mode_server(self):
sa = self._resolve(warmup_mode="server", explicit=("warmup_mode",))
self.assertEqual(sa.warmup_mode, "server")
self.assertTrue(sa.warmup)
self.assertTrue(sa.server_warmup)
def test_explicit_mode_overrides_explicit_legacy(self):
sa = self._resolve(
warmup_mode="request",
warmup=True,
server_warmup=True,
explicit=("warmup_mode", "warmup", "server_warmup"),
)
self.assertEqual(sa.warmup_mode, "request")
self.assertTrue(sa.warmup)
self.assertFalse(sa.server_warmup)
def test_explicit_legacy_false_beats_defaulted_mode(self):
# serve defaults warmup_mode="server" (not explicit); `--warmup false` wins.
sa = self._resolve(
warmup_mode="server",
warmup=False,
server_warmup=False,
explicit=("warmup",),
)
self.assertEqual(sa.warmup_mode, "off")
self.assertFalse(sa.warmup)
self.assertFalse(sa.server_warmup)
def test_defaulted_mode_applies_without_legacy_flags(self):
# bare `sglang serve`: warmup_mode="server" defaulted, no legacy override.
sa = self._resolve(warmup_mode="server")
self.assertEqual(sa.warmup_mode, "server")
self.assertTrue(sa.warmup)
self.assertTrue(sa.server_warmup)
def test_legacy_only_maps_to_request(self):
sa = self._resolve(warmup_mode=None, warmup=True, explicit=("warmup",))
self.assertEqual(sa.warmup_mode, "request")
self.assertTrue(sa.warmup)
self.assertFalse(sa.server_warmup)
def test_resolutions_force_warmup_on(self):
sa = self._resolve(
warmup_mode="off",
warmup_resolutions=["512x512"],
explicit=("warmup_mode",),
)
self.assertTrue(sa.warmup)
self.assertFalse(sa.server_warmup)
self.assertEqual(sa.warmup_mode, "request")
def test_disagg_role_disables_server_warmup(self):
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
sa = self._resolve(
warmup_mode="server",
disagg_role=RoleType.DENOISER,
explicit=("warmup_mode",),
)
self.assertTrue(sa.warmup)
self.assertFalse(sa.server_warmup)
self.assertEqual(sa.warmup_mode, "request")
def test_invalid_mode_raises(self):
with self.assertRaises(ValueError):
self._resolve(warmup_mode="bogus", explicit=("warmup_mode",))
class TestWarmupImageIsModelValid(unittest.TestCase):
"""The server-warmup placeholder image must be large enough for real pipelines."""
def test_minimum_warmup_image_is_at_least_64px(self):
import base64
import struct
from sglang.multimodal_gen.runtime.server_warmup import (
MINIMUM_PICTURE_BASE64_FOR_WARMUP,
)
payload = MINIMUM_PICTURE_BASE64_FOR_WARMUP.split(",", 1)[-1]
raw = base64.b64decode(payload)
self.assertEqual(raw[:8], b"\x89PNG\r\n\x1a\n")
# IHDR width/height are the two big-endian uint32 after the chunk header.
width, height = struct.unpack(">II", raw[16:24])
self.assertGreaterEqual(width, 64)
self.assertGreaterEqual(height, 64)
class TestOffloadDefaults(unittest.TestCase):
def _from_dict_with_pipeline_config(
self,