[diffusion] feat: add --warmup-mode enum server arg (#28184)
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user