diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py index b80ce7ed4..30c5dfe6e 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py @@ -36,6 +36,10 @@ from sglang.multimodal_gen.runtime.pipelines_core import Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch from sglang.multimodal_gen.runtime.scheduler_client import sync_scheduler_client from sglang.multimodal_gen.runtime.server_args import PortArgs, ServerArgs +from sglang.multimodal_gen.runtime.server_warmup import ( + run_sync_client_warmup, + should_run_explicit_client_warmup, +) from sglang.multimodal_gen.runtime.utils.logging_utils import ( GREEN, RESET, @@ -130,13 +134,13 @@ class DiffGenerator: logger.info(f"Local mode: {local_mode}") if local_mode: instance.local_scheduler_process = instance._start_local_server_if_needed() + instance.owns_scheduler_client = True + instance._run_client_warmup_if_needed() else: # In remote mode, we just need to connect and check. sync_scheduler_client.initialize(server_args) instance._check_remote_scheduler() - - # In both modes, this DiffGenerator instance is responsible for the client's lifecycle. - instance.owns_scheduler_client = True + instance.owns_scheduler_client = True return instance def _start_local_server_if_needed( @@ -150,6 +154,12 @@ class DiffGenerator: return processes + def _run_client_warmup_if_needed(self) -> None: + if not should_run_explicit_client_warmup(self.server_args): + return + + run_sync_client_warmup(self.server_args, sync_scheduler_client.forward) + def _check_remote_scheduler(self): """Check if the remote scheduler is accessible.""" if not sync_scheduler_client.ping(): diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py index 65911681f..c689580aa 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py @@ -35,12 +35,11 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import ( ) from sglang.multimodal_gen.runtime.scheduler_client import async_scheduler_client from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args -from sglang.multimodal_gen.runtime.server_warmup import prepare_warmup_image_path -from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger -from sglang.multimodal_gen.runtime.warmup_request_builder import ( - build_warmup_reqs, - should_include_warmup_image, +from sglang.multimodal_gen.runtime.server_warmup import ( + run_async_client_warmup, + should_run_synthetic_server_warmup, ) +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.srt.utils.json_response import orjson_response from sglang.version import __version__ @@ -73,56 +72,21 @@ async def _wait_until_http_ready(server_args: ServerArgs) -> None: raise RuntimeError(f"HTTP server did not become ready at {health_url}") -def _is_realtime_serving(server_args: ServerArgs) -> bool: - """A realtime pipeline establishes per-session state over the WebSocket, so - the synthetic server-warmup request (which has no session) cannot run — it - would fail in the realtime stage and abort startup. Detect it via the - realtime-adapter registry and skip server warmup.""" - try: - from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.registry import ( - get_realtime_model_adapter, - ) - - get_realtime_model_adapter(server_args) - return True - except Exception: - return False - - async def _run_server_warmup_after_http_ready( server_args: ServerArgs, warmup_done: asyncio.Event ) -> None: try: - if ( - not server_args.warmup - or not server_args.server_warmup - or server_args.warmup_resolutions is not None - or _is_realtime_serving(server_args) - ): + if not should_run_synthetic_server_warmup(server_args): warmup_done.set() return await _wait_until_http_ready(server_args) - warmup_input_path = None - if should_include_warmup_image(server_args, server_based_warmup=True): - warmup_input_path = await prepare_warmup_image_path(server_args) - - warmup_reqs = build_warmup_reqs( + await run_async_client_warmup( server_args, - warmup_resolutions=None, - warmup_input_path=warmup_input_path, - return_warmup_result=True, - server_based_warmup=True, - use_model_sampling_defaults=True, + async_scheduler_client.forward, + fail_open=server_args.warmup_resolutions is None, ) - warmup_total = len(warmup_reqs) - for req in warmup_reqs: - req.extra["warmup_total"] = warmup_total - response = await async_scheduler_client.forward(req) - if response.error is not None: - raise RuntimeError(response.error) - logger.info("The server is fired up and ready to roll!") warmup_done.set() except asyncio.CancelledError: diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py index f304f7ed5..9736791cc 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/utils.py @@ -1,10 +1,8 @@ # Copied and adapted from: https://github.com/hao-ai-lab/FastVideo import asyncio -import base64 import inspect import json import os -import re import shutil import tempfile import time @@ -30,6 +28,8 @@ from sglang.multimodal_gen.runtime.entrypoints.utils import ( from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch from sglang.multimodal_gen.runtime.scheduler_client import AsyncSchedulerClient from sglang.multimodal_gen.runtime.server_args import get_global_server_args +from sglang.multimodal_gen.runtime.utils.common import parse_size +from sglang.multimodal_gen.runtime.utils.image_io import save_base64_image_to_path from sglang.multimodal_gen.runtime.utils.logging_utils import ( init_logger, log_batch_completion, @@ -95,17 +95,6 @@ def temp_dir_if_disabled( shutil.rmtree(tmp, ignore_errors=True) -def _parse_size(size: str) -> tuple[int, int] | tuple[None, None]: - try: - parts = size.lower().replace(" ", "").split("x") - if len(parts) != 2: - raise ValueError - w, h = int(parts[0]), int(parts[1]) - return w, h - except Exception: - return None, None - - def choose_output_image_ext( output_format: Optional[str], background: Optional[str] ) -> str: @@ -135,7 +124,7 @@ def build_sampling_params(request_id: str, **kwargs) -> SamplingParams: # parse "WxH" size string if provided size = kwargs.pop("size", None) if size: - w, h = _parse_size(size) + w, h = parse_size(size) if w is not None: # treat None dimensions as unset so parsed size can fill them if kwargs.get("width") is None: @@ -227,7 +216,7 @@ async def _maybe_url_image( if prefer_remote_source: return img_url # encode image base64 url and persist on disk - input_path = await _save_base64_image_to_path(img_url, target_path) + input_path = save_base64_image_to_path(img_url, target_path) return input_path else: raise ValueError("Unsupported image url format") @@ -324,46 +313,6 @@ async def _save_url_image_to_path(image_url: str, target_path: str) -> str: ) -async def _save_base64_image_to_path(base64_data: str, target_path: str) -> str: - """Decode base64 image data and save to target path.""" - - _B64_FMT_HINT = ( - "Failed to decode base64 image. " - "Expected format: `data:[];base64,`" - ) - - # split `data:[][;base64],` to media-type base64 data - pattern = r"data:(.*?)(;base64)?,(.*)" - match = re.match(pattern, base64_data) - if not match: - raise ValueError(_B64_FMT_HINT) - media_type = match.group(1) - is_base64 = match.group(2) - if not is_base64: - raise ValueError(f"{_B64_FMT_HINT} (missing ;base64 marker)") - data = match.group(3) - if not data: - raise ValueError(f"{_B64_FMT_HINT} (empty data payload)") - # get ext from url - if media_type.startswith("image/"): - ext = media_type.split("/")[-1].lower() - if ext == "jpeg": - ext = "jpg" - else: - ext = "jpg" - target_path = f"{target_path}.{ext}" - os.makedirs(os.path.dirname(target_path), exist_ok=True) - - try: - image_data = base64.b64decode(data) - with open(target_path, "wb") as f: - f.write(image_data) - - return target_path - except Exception as e: - raise Exception(f"Failed to decode base64 image: {str(e)}") - - async def process_generation_batch( scheduler_client: AsyncSchedulerClient, batch, diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index 12bff0a24..752c9c37c 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -11,13 +11,11 @@ from enum import Enum from typing import Any, Iterator, List import zmq -from tqdm.auto import tqdm from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType from sglang.multimodal_gen.runtime.disaggregation.scheduler_mixin import ( SchedulerDisaggMixin, ) -from sglang.multimodal_gen.runtime.distributed import get_world_group from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import ( GetWeightsChecksumReqInput, UpdateWeightFromDiskReqInput, @@ -56,20 +54,14 @@ from sglang.multimodal_gen.runtime.server_args import ( set_global_server_args, ) from sglang.multimodal_gen.runtime.server_warmup import ( - get_first_generation_req, - is_server_based_warmup, + SchedulerWarmupMixin, is_warmup_req, - prepare_warmup_image_path_sync, should_return_warmup_result, ) from sglang.multimodal_gen.runtime.utils.common import get_zmq_socket from sglang.multimodal_gen.runtime.utils.distributed import broadcast_pyobj from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.trace_wrapper import DiffStage, trace_slice -from sglang.multimodal_gen.runtime.warmup_request_builder import ( - build_warmup_reqs, - should_include_warmup_image, -) logger = init_logger(__name__) @@ -77,7 +69,7 @@ _MAX_RECV_REQS_PER_POLL = 1024 _BATCH_METRICS_LOG_INTERVAL = 5 -class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin): +class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisaggMixin): """ Runs the main event loop for the rank 0 worker. It listens for external requests via ZMQ and coordinates with other workers. @@ -159,16 +151,13 @@ class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin): if self.receiver is not None: self._poller.register(self.receiver, zmq.POLLIN) - # whether we've send the necessary warmup reqs - self.warmed_up = False + self.req_based_warmup_scheduled = False # warmup progress tracking self._warmup_total = 0 self._warmup_processed = 0 self._warmup_progress_bar: Any | None = None self._logged_server_ready_after_warmup = False - self.prepare_server_warmup_reqs() - # Maximum consecutive errors before terminating the event loop self._max_consecutive_errors = 3 self._consecutive_error_count = 0 @@ -255,103 +244,6 @@ class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin): return [self._dispatch_single_request(req) for req in reqs] return self._dispatch_single_request(reqs[0]) - @staticmethod - def _format_warmup_req(req_or_group: Any) -> str: - req = get_first_generation_req(req_or_group) - prefix = ( - "server warmup req" - if is_server_based_warmup(req_or_group) - else "warmup req" - ) - if req is None: - return prefix - - shape = f"{req.width}x{req.height}" - if req.num_frames is not None and req.num_frames > 1: - shape += f"x{req.num_frames}f" - - default_steps = req.extra.get("cache_dit_num_inference_steps") - if default_steps is not None and default_steps != req.num_inference_steps: - steps = f"{req.num_inference_steps}/{default_steps} steps" - else: - steps = f"{req.num_inference_steps} step" - if req.num_inference_steps != 1: - steps += "s" - - return f"{prefix} ({shape}, {steps})" - - def _warmup_progress_total(self, req_or_group: Any | None = None) -> int: - req = get_first_generation_req(req_or_group) - if req is not None: - warmup_total = req.extra.get("warmup_total") - if warmup_total is not None: - return warmup_total - - return max(self._warmup_total, 1) - - def _ensure_warmup_progress_bar(self, req_or_group: Any) -> None: - if not self._show_warmup_progress: - return - - if self._warmup_progress_bar is None: - self._warmup_progress_bar = tqdm( - total=self._warmup_progress_total(req_or_group), - desc="Warmup requests", - unit="req", - ) - self._warmup_progress_bar.set_postfix_str( - self._format_warmup_req(req_or_group), refresh=False - ) - - def _advance_warmup_progress_bar( - self, req_or_group: Any, output_batch: OutputBatch - ) -> None: - if not self._show_warmup_progress: - return - - if self._warmup_progress_bar is None: - self._ensure_warmup_progress_bar(req_or_group) - - if output_batch.metrics is not None: - last_duration_s = output_batch.metrics.total_duration_s - self._warmup_progress_bar.set_postfix_str( - f"{self._format_warmup_req(req_or_group)}, last={last_duration_s:.2f}s", - refresh=False, - ) - self._warmup_progress_bar.update(1) - - if self._warmup_progress_bar.n >= self._warmup_progress_bar.total: - self._warmup_progress_bar.close() - self._warmup_progress_bar = None - - def _log_warmup_result( - self, - output_batch: OutputBatch, - req_or_group: Any, - is_warmup: bool, - ) -> None: - if not is_warmup: - return - - server_based_warmup = is_server_based_warmup(req_or_group) - self._warmup_processed += 1 - self._advance_warmup_progress_bar(req_or_group, output_batch) - - if output_batch.error is None: - if ( - not server_based_warmup - and not self._logged_server_ready_after_warmup - and ( - self._warmup_total <= 0 - or self._warmup_processed >= self._warmup_total - ) - ): - logger.info("The server is fired up and ready to roll!") - self._logged_server_ready_after_warmup = True - else: - warmup_desc = self._format_warmup_req(req_or_group) - logger.info(f"{warmup_desc} processing failed") - def _handle_generation( self, reqs: list[Any], *, allow_dynamic_batching: bool = True ): @@ -964,90 +856,6 @@ class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin): ) return batch_items - def prepare_server_warmup_reqs(self): - if ( - not self.server_args.warmup - or self.warmed_up - or self.server_args.warmup_resolutions is None - ): - return - - self._warmup_total = len(self.server_args.warmup_resolutions) - self._warmup_processed = 0 - - warmup_input_path = None - if should_include_warmup_image(self.server_args, server_based_warmup=False): - warmup_input_path = self._prepare_shared_warmup_image_path() - - warmup_reqs = build_warmup_reqs( - self.server_args, - warmup_resolutions=self.server_args.warmup_resolutions, - warmup_input_path=warmup_input_path, - ) - for req in warmup_reqs: - self.waiting_queue.append((None, req, time.monotonic())) - - # if server is warmed-up, set this flag to avoid req-based warmup - self.warmed_up = True - - def _prepare_shared_warmup_image_path(self) -> str: - world_group = get_world_group() - src_rank = world_group.ranks[0] - - warmup_sync: dict[str, str | None] - if world_group.rank == src_rank: - try: - input_path = prepare_warmup_image_path_sync(self.server_args) - warmup_sync = {"input_path": input_path, "error": None} - except Exception as e: - warmup_sync = {"input_path": None, "error": str(e)} - else: - warmup_sync = {} - - # Sync rank 0's warmup-image write result (path or error) to all ranks. - warmup_sync = broadcast_pyobj( - warmup_sync, - world_group.rank, - world_group.cpu_group, - src=src_rank, - ) - if not isinstance(warmup_sync, dict): - raise RuntimeError("Invalid warmup sync payload received across ranks") - - error = warmup_sync.get("error") - if error is not None: - raise RuntimeError( - f"Warmup image preparation failed on rank {src_rank}: {error}" - ) - - input_path = warmup_sync.get("input_path") - if not isinstance(input_path, str) or not input_path: - raise RuntimeError("Warmup image preparation returned empty input path") - - return input_path - - def process_received_reqs_with_req_based_warmup( - self, recv_reqs: List[tuple[bytes, Any]] - ) -> List[tuple[bytes, Any]]: - if ( - self.warmed_up - or not self.server_args.warmup - or not recv_reqs - or self.server_args.warmup_resolutions is not None - or self.server_args.server_warmup - ): - return recv_reqs - - identity, req_or_group = recv_reqs[0] - req = get_first_generation_req(req_or_group) - if req is not None: - warmup_req = req.copy_as_warmup(self.server_args.warmup_steps) - recv_reqs.insert(0, (identity, warmup_req)) - self._warmup_total = 1 - self._warmup_processed = 0 - self.warmed_up = True - return recv_reqs - @staticmethod def _normalize_received_payload( identity: bytes, reqs: Any diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py index 30b48a03f..66076500e 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py @@ -67,6 +67,7 @@ class PipelineStage(StageDedupMixin, ABC): # Class-level default so subclasses that override __init__ without # calling super().__init__() still see a consistent explicit-range gate. _current_use_nvtx: bool = False + _current_batch_is_warmup: bool = False def __init__(self): self.server_args = get_global_server_args() @@ -76,7 +77,7 @@ class PipelineStage(StageDedupMixin, ABC): def log_info(self, msg, *args): """Logs an informational message with the stage name as a prefix.""" - if self.server_args.comfyui_mode: + if self.server_args.comfyui_mode or self._current_batch_is_warmup: return logger.info(f"[{self.__class__.__name__}] {msg}", *args) @@ -355,6 +356,8 @@ class PipelineStage(StageDedupMixin, ABC): self._apply_nvtx_gate(batch.is_warmup) # Execute the actual stage logic with unified profiling. + previous_batch_is_warmup = self._current_batch_is_warmup + self._current_batch_is_warmup = batch.is_warmup try: with StageProfiler( stage_name, @@ -366,6 +369,7 @@ class PipelineStage(StageDedupMixin, ABC): ): result = self.forward(batch, server_args) finally: + self._current_batch_is_warmup = previous_batch_is_warmup self._current_use_nvtx = False # Post-execution output verification diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index 071f1c314..873aea97f 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -691,7 +691,6 @@ class ServerArgs(DisaggServerArgsMixin): def _adjust_warmup(self): if self.warmup_resolutions is not None: self.warmup = True - self.server_warmup = False if self.disagg_role != RoleType.MONOLITHIC: self.server_warmup = False @@ -1258,12 +1257,13 @@ class ServerArgs(DisaggServerArgsMixin): default=ServerArgs.warmup, help=( "Perform warmup before normal traffic. `sglang serve` runs a " - "lightweight server warmup after HTTP is ready; other entrypoints " - "use request-based warmup unless `--warmup-resolutions` is " - "specified. Recommended to enable when benchmarking to ensure fair " - "comparison and best performance. When enabled with " - "`--warmup-resolutions` unspecified, look for the line ending with " - "`(with warmup excluded)` for actual processing time." + "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." ), ) parser.add_argument( @@ -1271,7 +1271,7 @@ class ServerArgs(DisaggServerArgsMixin): type=str, nargs="+", default=ServerArgs.warmup_resolutions, - help="Specify resolutions for server to warmup. e.g., `--warmup-resolutions 256x256, 720x720`", + help="Specify explicit warmup resolutions. e.g., `--warmup-resolutions 256x256 720x720`", ) parser.add_argument( "--warmup-steps", diff --git a/python/sglang/multimodal_gen/runtime/server_warmup.py b/python/sglang/multimodal_gen/runtime/server_warmup.py index d6d21ce70..a0243d600 100644 --- a/python/sglang/multimodal_gen/runtime/server_warmup.py +++ b/python/sglang/multimodal_gen/runtime/server_warmup.py @@ -1,15 +1,23 @@ # Copied and adapted from: https://github.com/hao-ai-lab/FastVideo # SPDX-License-Identifier: Apache-2.0 -import asyncio import os import tempfile -from typing import Any +from typing import Any, Awaitable, Callable -from sglang.multimodal_gen.runtime.entrypoints.openai.utils import save_image_to_path -from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req +from tqdm.auto import tqdm + +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( + OutputBatch, + Req, +) from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.utils.image_io import save_base64_image_to_path from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.warmup_request_builder import ( + build_warmup_reqs, + should_include_warmup_image, +) logger = init_logger(__name__) @@ -50,7 +58,115 @@ def should_return_warmup_result(req_or_group: Any) -> bool: ) -async def prepare_warmup_image_path(server_args: ServerArgs) -> str: +def should_run_server_warmup(server_args: ServerArgs) -> bool: + return server_args.warmup and server_args.server_warmup + + +def is_realtime_serving(server_args: ServerArgs) -> bool: + """Synthetic warmup has no realtime session state.""" + try: + from sglang.multimodal_gen.runtime.entrypoints.openai.realtime.registry import ( + get_realtime_model_adapter, + ) + + get_realtime_model_adapter(server_args) + return True + except Exception: + return False + + +def should_run_synthetic_server_warmup(server_args: ServerArgs) -> bool: + return should_run_server_warmup(server_args) and not is_realtime_serving( + server_args + ) + + +def should_run_explicit_client_warmup(server_args: ServerArgs) -> bool: + return server_args.warmup and server_args.warmup_resolutions is not None + + +def format_warmup_req(req_or_group: Any) -> str: + req = get_first_generation_req(req_or_group) + prefix = ( + "server warmup req" if is_server_based_warmup(req_or_group) else "warmup req" + ) + if req is None: + return prefix + + shape = f"{req.width}x{req.height}" + if req.num_frames is not None and req.num_frames > 1: + shape += f"x{req.num_frames}f" + + default_steps = req.extra.get("cache_dit_num_inference_steps") + if default_steps is not None and default_steps != req.num_inference_steps: + steps = f"{req.num_inference_steps}/{default_steps} steps" + else: + steps = f"{req.num_inference_steps} step" + if req.num_inference_steps != 1: + steps += "s" + + return f"{prefix} ({shape}, {steps})" + + +def build_client_warmup_reqs( + server_args: ServerArgs, + *, + warmup_input_path: str | None = None, +) -> list[Req]: + warmup_reqs = build_warmup_reqs( + server_args, + warmup_resolutions=server_args.warmup_resolutions, + warmup_input_path=warmup_input_path, + return_warmup_result=True, + server_based_warmup=True, + ) + warmup_total = len(warmup_reqs) + for req in warmup_reqs: + req.extra["warmup_total"] = warmup_total + return warmup_reqs + + +async def run_async_client_warmup( + server_args: ServerArgs, + forward: Callable[[Req], Awaitable[OutputBatch]], + *, + fail_open: bool = False, +) -> None: + try: + warmup_input_path = None + if should_include_warmup_image(server_args, server_based_warmup=True): + warmup_input_path = prepare_warmup_image_path(server_args) + + for req in build_client_warmup_reqs( + server_args, warmup_input_path=warmup_input_path + ): + response = await forward(req) + if response.error is not None: + raise RuntimeError(response.error) + except Exception as e: + if fail_open: + logger.warning("Synthetic server warmup failed; continuing startup: %s", e) + return + raise + + +def run_sync_client_warmup( + server_args: ServerArgs, + forward: Callable[[Req], OutputBatch], +) -> None: + warmup_input_path = None + if should_include_warmup_image(server_args, server_based_warmup=True): + warmup_input_path = prepare_warmup_image_path(server_args) + + for req in build_client_warmup_reqs( + server_args, warmup_input_path=warmup_input_path + ): + response = forward(req) + if response.error is not None: + raise RuntimeError(response.error) + + +def prepare_warmup_image_path(server_args: ServerArgs) -> str: if server_args.input_save_path is not None: uploads_dir = server_args.input_save_path os.makedirs(uploads_dir, exist_ok=True) @@ -58,10 +174,106 @@ async def prepare_warmup_image_path(server_args: ServerArgs) -> str: uploads_dir = tempfile.mkdtemp(prefix="sglang_input_") warmup_image_base = os.path.join(uploads_dir, "warmup_image") - return await save_image_to_path( + return save_base64_image_to_path( MINIMUM_PICTURE_BASE64_FOR_WARMUP, warmup_image_base ) -def prepare_warmup_image_path_sync(server_args: ServerArgs) -> str: - return asyncio.run(prepare_warmup_image_path(server_args)) +class SchedulerWarmupMixin: + @staticmethod + def _format_warmup_req(req_or_group: Any) -> str: + return format_warmup_req(req_or_group) + + def _warmup_progress_total(self, req_or_group: Any | None = None) -> int: + req = get_first_generation_req(req_or_group) + if req is not None: + warmup_total = req.extra.get("warmup_total") + if warmup_total is not None: + return warmup_total + + return max(self._warmup_total, 1) + + def _ensure_warmup_progress_bar(self, req_or_group: Any) -> None: + if not self._show_warmup_progress: + return + + if self._warmup_progress_bar is None: + self._warmup_progress_bar = tqdm( + total=self._warmup_progress_total(req_or_group), + desc="Warmup requests", + unit="req", + ) + self._warmup_progress_bar.set_postfix_str( + self._format_warmup_req(req_or_group), refresh=False + ) + + def _advance_warmup_progress_bar( + self, req_or_group: Any, output_batch: OutputBatch + ) -> None: + if not self._show_warmup_progress: + return + + if self._warmup_progress_bar is None: + self._ensure_warmup_progress_bar(req_or_group) + + if output_batch.metrics is not None: + last_duration_s = output_batch.metrics.total_duration_s + self._warmup_progress_bar.set_postfix_str( + f"{self._format_warmup_req(req_or_group)}, last={last_duration_s:.2f}s", + refresh=False, + ) + self._warmup_progress_bar.update(1) + + if self._warmup_progress_bar.n >= self._warmup_progress_bar.total: + self._warmup_progress_bar.close() + self._warmup_progress_bar = None + + def _log_warmup_result( + self, + output_batch: OutputBatch, + req_or_group: Any, + is_warmup: bool, + ) -> None: + if not is_warmup: + return + + server_based_warmup = is_server_based_warmup(req_or_group) + self._warmup_processed += 1 + self._advance_warmup_progress_bar(req_or_group, output_batch) + + if output_batch.error is None: + if ( + not server_based_warmup + and not self._logged_server_ready_after_warmup + and ( + self._warmup_total <= 0 + or self._warmup_processed >= self._warmup_total + ) + ): + logger.info("The server is fired up and ready to roll!") + self._logged_server_ready_after_warmup = True + else: + warmup_desc = self._format_warmup_req(req_or_group) + logger.info(f"{warmup_desc} processing failed") + + def process_received_reqs_with_req_based_warmup( + self, recv_reqs: list[tuple[bytes, Any]] + ) -> list[tuple[bytes, Any]]: + if ( + self.req_based_warmup_scheduled + or not self.server_args.warmup + or not recv_reqs + or self.server_args.warmup_resolutions is not None + or self.server_args.server_warmup + ): + return recv_reqs + + identity, req_or_group = recv_reqs[0] + req = get_first_generation_req(req_or_group) + if req is not None: + warmup_req = req.copy_as_warmup(self.server_args.warmup_steps) + recv_reqs.insert(0, (identity, warmup_req)) + self._warmup_total = 1 + self._warmup_processed = 0 + self.req_based_warmup_scheduled = True + return recv_reqs diff --git a/python/sglang/multimodal_gen/runtime/utils/common.py b/python/sglang/multimodal_gen/runtime/utils/common.py index 6e9459ce9..a921130b5 100644 --- a/python/sglang/multimodal_gen/runtime/utils/common.py +++ b/python/sglang/multimodal_gen/runtime/utils/common.py @@ -110,6 +110,16 @@ def normalize_gpu_ids(gpu_ids: Any) -> list[int] | None: return parsed +def parse_size(size: str) -> tuple[int | None, int | None]: + try: + parts = size.lower().replace(" ", "").split("x") + if len(parts) != 2: + raise ValueError + return int(parts[0]), int(parts[1]) + except ValueError: + return None, None + + def parse_tcp_host_port(value: str | None, field_name: str) -> tuple[str, int]: if value is None or not str(value).strip(): raise ValueError(f"{field_name} is required") diff --git a/python/sglang/multimodal_gen/runtime/utils/image_io.py b/python/sglang/multimodal_gen/runtime/utils/image_io.py new file mode 100644 index 000000000..e40b64915 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/utils/image_io.py @@ -0,0 +1,41 @@ +# SPDX-License-Identifier: Apache-2.0 +import base64 +import os +import re + + +def save_base64_image_to_path(base64_data: str, target_path: str) -> str: + b64_format_hint = ( + "Failed to decode base64 image. " + "Expected format: `data:[];base64,`" + ) + + match = re.match(r"data:(.*?)(;base64)?,(.*)", base64_data) + if not match: + raise ValueError(b64_format_hint) + media_type = match.group(1) + is_base64 = match.group(2) + if not is_base64: + raise ValueError(f"{b64_format_hint} (missing ;base64 marker)") + data = match.group(3) + if not data: + raise ValueError(f"{b64_format_hint} (empty data payload)") + + if media_type.startswith("image/"): + ext = media_type.split("/")[-1].lower() + if ext == "jpeg": + ext = "jpg" + else: + ext = "jpg" + target_path = f"{target_path}.{ext}" + os.makedirs(os.path.dirname(target_path), exist_ok=True) + + try: + image_data = base64.b64decode(data) + except Exception as exc: + raise Exception(f"Failed to decode base64 image: {str(exc)}") from exc + + with open(target_path, "wb") as f: + f.write(image_data) + + return target_path diff --git a/python/sglang/multimodal_gen/runtime/warmup_request_builder.py b/python/sglang/multimodal_gen/runtime/warmup_request_builder.py index 046dcdeb8..106a6822b 100644 --- a/python/sglang/multimodal_gen/runtime/warmup_request_builder.py +++ b/python/sglang/multimodal_gen/runtime/warmup_request_builder.py @@ -1,58 +1,74 @@ # SPDX-License-Identifier: Apache-2.0 """Build synthetic diffusion warmup requests. -Default server warmup should cover the model's normal serving path before the -first real request, without copying user traffic. It therefore starts from the -model's sampling defaults, keeps default resolution/frame semantics, and trims -only the denoising step count. +Default server warmup should cover a representative serving path before the +first real request, without copying user traffic. It starts from the model's +sampling defaults, then keeps startup bounded by choosing common low-cost +resolution/frame buckets and trimming the denoising step count. Image models may run a tiny second step because first/last step paths often -initialize different kernels or scheduler state. Video models stay at -`warmup_steps` to keep startup bounded. Explicit request-based warmup remains a -scheduler-level legacy path and is not constructed here. +initialize different kernels or scheduler state. Video models cap frames and +steps to keep startup bounded. Explicit warmup resolutions share this builder; +callers send them through the scheduler client so warmup exercises the same +request transport path as real generation. """ from copy import copy from typing import Any from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType -from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams +from sglang.multimodal_gen.configs.sample.sampling_params import ( + DataType, + SamplingParams, +) from sglang.multimodal_gen.registry import get_pipeline_config_classes -from sglang.multimodal_gen.runtime.entrypoints.openai.utils import _parse_size from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req -from sglang.multimodal_gen.runtime.server_args import ServerArgs +from sglang.multimodal_gen.runtime.server_args import ( + ServerArgs, + is_ltx2_two_stage_pipeline_name, +) +from sglang.multimodal_gen.runtime.utils.common import parse_size from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) DEFAULT_PLACEHOLDER_PROMPT = "warmup" DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION = (64, 64) +SERVER_WARMUP_IMAGE_FALLBACK_RESOLUTION = (512, 512) +SERVER_WARMUP_VIDEO_FALLBACK_RESOLUTION = (832, 480) +SERVER_WARMUP_IMAGE_MAX_AREA = 768 * 768 +SERVER_WARMUP_DIFFUSERS_IMAGE_MAX_AREA = 512 * 512 +SERVER_WARMUP_VIDEO_MAX_AREA = 832 * 480 +SERVER_WARMUP_MAX_VIDEO_FRAMES = 17 SERVER_WARMUP_IMAGE_STEPS = 2 +SERVER_WARMUP_VIDEO_STEPS = 2 def get_model_sampling_defaults(server_args: ServerArgs) -> SamplingParams: pipeline_class_name = server_args.pipeline_class_name - try: - if pipeline_class_name: - config_classes = get_pipeline_config_classes(pipeline_class_name) - if config_classes is not None: - _, sampling_params_cls = config_classes - return sampling_params_cls() + if pipeline_class_name: + config_classes = get_pipeline_config_classes(pipeline_class_name) + if config_classes is not None: + _, sampling_params_cls = config_classes + return sampling_params_cls() - return SamplingParams.from_pretrained( - server_args.model_path, - backend=server_args.backend, - model_id=server_args.model_id, - ) - except Exception: - logger.debug("Falling back to base SamplingParams for server warmup") - return SamplingParams() + return SamplingParams.from_pretrained( + server_args.model_path, + backend=server_args.backend, + model_id=server_args.model_id, + ) def _resolve_default_warmup_resolution( server_args: ServerArgs, sampling_defaults: SamplingParams, + *, + server_based_warmup: bool, ) -> tuple[int, int]: + """returns a default resolution to warmup""" + if server_based_warmup: + return _resolve_representative_warmup_resolution(server_args, sampling_defaults) + width = sampling_defaults.width height = sampling_defaults.height if width is not None and height is not None: @@ -71,6 +87,150 @@ def _resolve_default_warmup_resolution( ) +def _resolve_representative_warmup_resolution( + server_args: ServerArgs, + sampling_defaults: SamplingParams, +) -> tuple[int, int]: + target_area = _target_warmup_area(server_args) + alignment = _warmup_resolution_alignment(server_args) + + supported_resolution = _select_supported_warmup_resolution( + sampling_defaults.supported_resolutions, target_area, alignment + ) + if supported_resolution is not None: + return supported_resolution + + width = sampling_defaults.width + height = sampling_defaults.height + if width is not None and height is not None: + return _fit_resolution_to_area(width, height, target_area, alignment) + + width, height = _fallback_warmup_resolution(server_args) + return _fit_resolution_to_area(width, height, target_area, alignment) + + +def _target_warmup_area(server_args: ServerArgs) -> int: + if server_args.pipeline_config.task_type.is_image_gen(): + if getattr(server_args, "backend", None) == "diffusers": + return SERVER_WARMUP_DIFFUSERS_IMAGE_MAX_AREA + return SERVER_WARMUP_IMAGE_MAX_AREA + if _is_video_warmup_task(server_args): + return SERVER_WARMUP_VIDEO_MAX_AREA + return ( + DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION[0] + * DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION[1] + ) + + +def _fallback_warmup_resolution(server_args: ServerArgs) -> tuple[int, int]: + if server_args.pipeline_config.task_type.is_image_gen(): + return SERVER_WARMUP_IMAGE_FALLBACK_RESOLUTION + if _is_video_warmup_task(server_args): + return SERVER_WARMUP_VIDEO_FALLBACK_RESOLUTION + return DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION + + +def _is_video_warmup_task(server_args: ServerArgs) -> bool: + return server_args.pipeline_config.task_type.data_type() == DataType.VIDEO + + +def _warmup_resolution_alignment(server_args: ServerArgs) -> int: + pipeline_config = server_args.pipeline_config + alignment = 16 + + vae_stride = getattr(pipeline_config, "vae_stride", None) + if vae_stride is not None: + spatial_stride = ( + vae_stride[-2:] if isinstance(vae_stride, (tuple, list)) else (vae_stride,) + ) + for stride in spatial_stride: + alignment = max(alignment, int(stride)) + + vae_scale_factor = getattr(pipeline_config, "vae_scale_factor", None) + if vae_scale_factor is not None: + alignment = max(alignment, int(vae_scale_factor)) + + arch_config = getattr( + getattr(pipeline_config, "vae_config", None), "arch_config", None + ) + for attr in ("vae_scale_factor", "spatial_compression_ratio"): + value = getattr(arch_config, attr, None) + if value is not None: + alignment = max(alignment, int(value)) + + if is_ltx2_two_stage_pipeline_name(server_args.pipeline_class_name): + vae_scale_factor = pipeline_config.vae_scale_factor + alignment = max(alignment, 64, int(vae_scale_factor) * 2) + + return alignment + + +def _select_supported_warmup_resolution( + supported_resolutions: list[tuple[int, int]] | None, + target_area: int, + alignment: int, +) -> tuple[int, int] | None: + if not supported_resolutions: + return None + + candidates = [ + resolution + for resolution in supported_resolutions + if resolution[0] * resolution[1] <= target_area + and _is_resolution_aligned(resolution, alignment) + ] + if candidates: + return max(candidates, key=lambda size: size[0] * size[1]) + + aligned_resolutions = [ + resolution + for resolution in supported_resolutions + if _is_resolution_aligned(resolution, alignment) + ] + if aligned_resolutions: + return min(aligned_resolutions, key=lambda size: size[0] * size[1]) + return None + + +def _fit_resolution_to_area( + width: int, height: int, target_area: int, alignment: int +) -> tuple[int, int]: + """adjust the warmup resolution to balance between warmup time and warmup effect""" + area = width * height + if area > target_area: + scale = (target_area / area) ** 0.5 + width = int(width * scale) + height = int(height * scale) + + return ( + max(alignment, width // alignment * alignment), + max(alignment, height // alignment * alignment), + ) + + +def _is_resolution_aligned(resolution: tuple[int, int], alignment: int) -> bool: + width, height = resolution + return width % alignment == 0 and height % alignment == 0 + + +def _resolve_warmup_num_frames( + server_args: ServerArgs, + sampling_defaults: SamplingParams, + *, + server_based_warmup: bool, +) -> int: + num_frames = sampling_defaults.num_frames + if ( + not server_based_warmup + or not _is_video_warmup_task(server_args) + or num_frames is None + ): + # use default num frames + return num_frames + + return min(num_frames, SERVER_WARMUP_MAX_VIDEO_FRAMES) + + def _effective_cfg_scale(sampling_defaults: SamplingParams) -> float | None: if sampling_defaults.true_cfg_scale is not None: return sampling_defaults.true_cfg_scale @@ -82,20 +242,22 @@ def _resolve_warmup_steps( sampling_defaults: SamplingParams, *, server_based_warmup: bool, - use_model_sampling_defaults: bool, ) -> int: warmup_steps = server_args.warmup_steps - if not server_based_warmup or not use_model_sampling_defaults: + if not server_based_warmup: return warmup_steps default_steps = sampling_defaults.num_inference_steps if default_steps is None or default_steps <= warmup_steps: return warmup_steps - if not server_args.pipeline_config.task_type.is_image_gen(): - return warmup_steps + if _is_video_warmup_task(server_args): + return min(default_steps, max(warmup_steps, SERVER_WARMUP_VIDEO_STEPS)) - return min(default_steps, max(warmup_steps, SERVER_WARMUP_IMAGE_STEPS)) + if server_args.pipeline_config.task_type.is_image_gen(): + return min(default_steps, max(warmup_steps, SERVER_WARMUP_IMAGE_STEPS)) + + return warmup_steps def should_include_warmup_image( @@ -118,35 +280,34 @@ def build_warmup_reqs( warmup_input_path: str | None = None, return_warmup_result: bool = False, server_based_warmup: bool = False, - use_model_sampling_defaults: bool = False, ) -> list[Req]: task_type = server_args.pipeline_config.task_type - if warmup_resolutions is None or use_model_sampling_defaults: - sampling_defaults = get_model_sampling_defaults(server_args) - else: - sampling_defaults = SamplingParams() + sampling_defaults = get_model_sampling_defaults(server_args) if warmup_resolutions is None: width, height = _resolve_default_warmup_resolution( - server_args, sampling_defaults + server_args, + sampling_defaults, + server_based_warmup=server_based_warmup, ) resolutions: list[tuple[int, int]] = [(width, height)] else: - resolutions = [_parse_size(resolution) for resolution in warmup_resolutions] + resolutions = [parse_size(resolution) for resolution in warmup_resolutions] - negative_prompt: Any = ( - sampling_defaults.negative_prompt if use_model_sampling_defaults else None - ) - cfg_scale = ( - _effective_cfg_scale(sampling_defaults) if use_model_sampling_defaults else None - ) + negative_prompt: Any = sampling_defaults.negative_prompt + cfg_scale = _effective_cfg_scale(sampling_defaults) warmup_steps = _resolve_warmup_steps( server_args, sampling_defaults, server_based_warmup=server_based_warmup, - use_model_sampling_defaults=use_model_sampling_defaults, + ) + warmup_num_frames = _resolve_warmup_num_frames( + server_args, + sampling_defaults, + server_based_warmup=server_based_warmup, ) + # build warmup reqs warmup_reqs = [] include_warmup_image = should_include_warmup_image(server_args, server_based_warmup) for width, height in resolutions: @@ -156,37 +317,27 @@ def build_warmup_reqs( height=height, prompt=DEFAULT_PLACEHOLDER_PROMPT, ) - if use_model_sampling_defaults: - req_kwargs["sampling_params"] = copy(sampling_defaults) - req_kwargs.update( - negative_prompt=negative_prompt, - guidance_scale=sampling_defaults.guidance_scale, - guidance_scale_2=sampling_defaults.guidance_scale_2, - true_cfg_scale=sampling_defaults.true_cfg_scale, - num_inference_steps=sampling_defaults.num_inference_steps, - num_frames=sampling_defaults.num_frames, - ) + req_kwargs["sampling_params"] = copy(sampling_defaults) + req_kwargs.update( + negative_prompt=negative_prompt, + guidance_scale=sampling_defaults.guidance_scale, + guidance_scale_2=sampling_defaults.guidance_scale_2, + true_cfg_scale=sampling_defaults.true_cfg_scale, + num_inference_steps=sampling_defaults.num_inference_steps, + num_frames=warmup_num_frames, + ) if include_warmup_image: if warmup_input_path is None: raise RuntimeError( "Warmup image path is required for image-input model" ) req_kwargs["prompt"] = DEFAULT_PLACEHOLDER_PROMPT - if not use_model_sampling_defaults: - req_kwargs["negative_prompt"] = "" req_kwargs["image_path"] = [warmup_input_path] - if ( - server_args.enable_cfg_parallel - and req_kwargs.get("negative_prompt") is None - ): - req_kwargs["negative_prompt"] = DEFAULT_PLACEHOLDER_PROMPT + if server_args.enable_cfg_parallel: + if not req_kwargs.get("negative_prompt"): + req_kwargs["negative_prompt"] = DEFAULT_PLACEHOLDER_PROMPT req_kwargs["do_classifier_free_guidance"] = True - elif ( - use_model_sampling_defaults - and negative_prompt is not None - and cfg_scale is not None - and cfg_scale > 1.0 - ): + elif negative_prompt is not None and cfg_scale is not None and cfg_scale > 1.0: req_kwargs["do_classifier_free_guidance"] = True req = Req(**req_kwargs) diff --git a/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py b/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py index e0705ee85..00c3d0b81 100644 --- a/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py +++ b/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py @@ -1,13 +1,14 @@ """Unit tests for the --enable-cfg-parallel warmup fix and guard. Covers warmup and cfg-parallel guard paths introduced alongside this file: -- Scheduler.prepare_server_warmup_reqs synthesizes warmup Reqs that - actually enable classifier-free guidance when cfg-parallel is on. +- build_warmup_reqs synthesizes warmup Reqs that actually enable + classifier-free guidance when cfg-parallel is on. +- DiffGenerator sends explicit warmup resolutions through the scheduler client. - InputValidationStage.forward rejects non-CFG requests when the server has cfg-parallel on. - Server-based warmup can opt into model-default negative prompts so warmup populates the negative text embedding cache. -- Req-based warmup remains available only through the explicit legacy path. +- Req-based warmup remains available only through the lazy legacy path. All tests are CPU-only; no model loading, no distributed init. """ @@ -24,8 +25,12 @@ from sglang.multimodal_gen.configs.pipeline_configs.flux_finetuned import ( Flux2FinetunedPipelineConfig, ) from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams +from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import DiffGenerator from sglang.multimodal_gen.runtime.managers.scheduler import Scheduler -from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req +from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( + OutputBatch, + Req, +) from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import ( ImageVAEEncodingStage, ) @@ -33,8 +38,8 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import InputValidationStage, ) from sglang.multimodal_gen.runtime.warmup_request_builder import ( - DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION, DEFAULT_PLACEHOLDER_PROMPT, + SERVER_WARMUP_IMAGE_FALLBACK_RESOLUTION, build_warmup_reqs, should_include_warmup_image, ) @@ -44,8 +49,7 @@ def _make_bare_scheduler(enable_cfg_parallel: bool) -> Scheduler: """ Build a minimal Scheduler without calling __init__ (which requires distributed init, ZMQ sockets, pipeline load, etc.). Populates only - the attributes prepare_server_warmup_reqs reads/writes for a - text-only task so _prepare_shared_warmup_image_path is skipped. + the attributes req-based warmup reads/writes. """ scheduler = object.__new__(Scheduler) @@ -54,10 +58,8 @@ def _make_bare_scheduler(enable_cfg_parallel: bool) -> Scheduler: server_args.warmup_steps = 1 server_args.warmup_resolutions = ["512x512"] server_args.enable_cfg_parallel = enable_cfg_parallel + server_args.server_warmup = False - # Text-only task — accepts_image_input() False skips the image-path - # branch entirely, so we don't need to mock - # _prepare_shared_warmup_image_path. task_type = MagicMock() task_type.requires_image_input.return_value = False task_type.accepts_image_input.return_value = False @@ -65,7 +67,7 @@ def _make_bare_scheduler(enable_cfg_parallel: bool) -> Scheduler: server_args.pipeline_config.task_type = task_type scheduler.server_args = server_args - scheduler.warmed_up = False + scheduler.req_based_warmup_scheduled = False scheduler.waiting_queue = deque() return scheduler @@ -92,14 +94,34 @@ def _make_validation_server_args(enable_cfg_parallel: bool) -> MagicMock: class TestWarmupReqCfgParallel(unittest.TestCase): - """Commit 1 regression: prepare_server_warmup_reqs.""" + """Warmup request construction and req-based warmup guards.""" def test_warmup_req_cfg_parallel_sets_do_cfg(self): - scheduler = _make_bare_scheduler(enable_cfg_parallel=True) - scheduler.prepare_server_warmup_reqs() + server_args = _make_bare_scheduler(enable_cfg_parallel=True).server_args + sampling_defaults = SamplingParams() + with patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults", + return_value=sampling_defaults, + ): + req = build_warmup_reqs( + server_args, + warmup_resolutions=["512x512"], + server_based_warmup=True, + )[0] + self.assertIs(req.do_classifier_free_guidance, True) + self.assertEqual(req.negative_prompt, sampling_defaults.negative_prompt) - self.assertEqual(len(scheduler.waiting_queue), 1) - _, req, _ = scheduler.waiting_queue[0] + def test_warmup_req_cfg_parallel_fills_missing_negative_prompt(self): + server_args = _make_bare_scheduler(enable_cfg_parallel=True).server_args + with patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults", + return_value=SamplingParams(negative_prompt=None), + ): + req = build_warmup_reqs( + server_args, + warmup_resolutions=["512x512"], + server_based_warmup=True, + )[0] self.assertIs(req.do_classifier_free_guidance, True) self.assertEqual(req.negative_prompt, DEFAULT_PLACEHOLDER_PROMPT) @@ -109,11 +131,16 @@ class TestWarmupReqCfgParallel(unittest.TestCase): # AND the synthesized Req is not using the cfg-parallel-specific # "warmup" placeholder for negative_prompt (which would indicate # the fix's kwargs leaked into this branch). - scheduler = _make_bare_scheduler(enable_cfg_parallel=False) - scheduler.prepare_server_warmup_reqs() - - self.assertEqual(len(scheduler.waiting_queue), 1) - _, req, _ = scheduler.waiting_queue[0] + server_args = _make_bare_scheduler(enable_cfg_parallel=False).server_args + with patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults", + return_value=SamplingParams(), + ): + req = build_warmup_reqs( + server_args, + warmup_resolutions=["512x512"], + server_based_warmup=True, + )[0] self.assertIs(req.do_classifier_free_guidance, False) self.assertNotEqual(req.negative_prompt, DEFAULT_PLACEHOLDER_PROMPT) @@ -133,7 +160,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): self.assertTrue(processed[0][1].metrics.suppress_stage_breakdown) self.assertEqual(processed[0][1].num_inference_steps, 1) self.assertEqual(processed[0][1].extra["cache_dit_num_inference_steps"], 20) - self.assertTrue(scheduler.warmed_up) + self.assertTrue(scheduler.req_based_warmup_scheduled) def test_req_based_warmup_skips_default_server_warmup_path(self): scheduler = _make_bare_scheduler(enable_cfg_parallel=False) @@ -145,7 +172,46 @@ class TestWarmupReqCfgParallel(unittest.TestCase): self.assertIs(processed, recv_reqs) self.assertEqual(len(processed), 1) - self.assertFalse(scheduler.warmed_up) + self.assertFalse(scheduler.req_based_warmup_scheduled) + + def test_diff_generator_runs_explicit_warmup_through_scheduler_client(self): + generator = object.__new__(DiffGenerator) + server_args = MagicMock() + server_args.warmup = True + server_args.warmup_resolutions = ["832x480"] + server_args.warmup_steps = 1 + server_args.enable_cfg_parallel = False + + task_type = MagicMock() + task_type.requires_image_input.return_value = False + task_type.accepts_image_input.return_value = False + task_type.is_image_gen.return_value = False + task_type.data_type.return_value = ModelTaskType.T2V.data_type() + server_args.pipeline_config.task_type = task_type + generator.server_args = server_args + + sampling_defaults = SamplingParams(num_frames=81, num_inference_steps=50) + with ( + patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults", + return_value=sampling_defaults, + ), + patch( + "sglang.multimodal_gen.runtime.entrypoints.diffusion_generator.sync_scheduler_client.forward", + return_value=OutputBatch(error=None), + ) as forward, + ): + generator._run_client_warmup_if_needed() + + forward.assert_called_once() + req = forward.call_args.args[0] + self.assertTrue(req.is_warmup) + self.assertEqual((req.width, req.height), (832, 480)) + self.assertEqual(req.num_frames, 17) + self.assertEqual(req.num_inference_steps, 2) + self.assertTrue(req.extra["return_warmup_result"]) + self.assertTrue(req.extra["server_based_warmup"]) + self.assertEqual(req.extra["warmup_total"], 1) def test_server_based_warmup_uses_model_default_negative_prompt(self): server_args = MagicMock() @@ -171,7 +237,6 @@ class TestWarmupReqCfgParallel(unittest.TestCase): reqs = build_warmup_reqs( server_args, warmup_resolutions=None, - use_model_sampling_defaults=True, return_warmup_result=True, server_based_warmup=True, ) @@ -183,6 +248,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): self.assertEqual(req.num_inference_steps, 2) self.assertEqual(req.extra["cache_dit_num_inference_steps"], 20) self.assertEqual(req.negative_prompt, "model default negative") + self.assertEqual(req.prompt, DEFAULT_PLACEHOLDER_PROMPT) self.assertIs(req.do_classifier_free_guidance, True) self.assertTrue(req.extra["return_warmup_result"]) self.assertTrue(req.extra["server_based_warmup"]) @@ -207,7 +273,6 @@ class TestWarmupReqCfgParallel(unittest.TestCase): reqs = build_warmup_reqs( server_args, warmup_resolutions=None, - use_model_sampling_defaults=True, server_based_warmup=True, ) @@ -215,7 +280,44 @@ class TestWarmupReqCfgParallel(unittest.TestCase): self.assertEqual(req.width, 640) self.assertEqual(req.height, 640) - def test_server_based_warmup_prefers_default_resolution_over_supported_min(self): + def test_server_based_warmup_resolutions_keep_sampling_defaults_and_caps(self): + server_args = MagicMock() + server_args.warmup_steps = 1 + server_args.enable_cfg_parallel = False + + task_type = MagicMock() + task_type.requires_image_input.return_value = False + task_type.accepts_image_input.return_value = False + task_type.is_image_gen.return_value = False + task_type.data_type.return_value = ModelTaskType.T2V.data_type() + server_args.pipeline_config.task_type = task_type + + sampling_defaults = SamplingParams( + negative_prompt="model default negative", + guidance_scale=3.5, + num_frames=81, + num_inference_steps=50, + ) + with patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults", + return_value=sampling_defaults, + ): + reqs = build_warmup_reqs( + server_args, + warmup_resolutions=["832x480"], + return_warmup_result=True, + server_based_warmup=True, + ) + + req = reqs[0] + self.assertEqual((req.width, req.height), (832, 480)) + self.assertEqual(req.num_frames, 17) + self.assertEqual(req.num_inference_steps, 2) + self.assertEqual(req.extra["cache_dit_num_inference_steps"], 50) + self.assertEqual(req.negative_prompt, "model default negative") + self.assertIs(req.do_classifier_free_guidance, True) + + def test_server_based_warmup_uses_supported_resolution_within_budget(self): server_args = MagicMock() server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False @@ -239,11 +341,62 @@ class TestWarmupReqCfgParallel(unittest.TestCase): reqs = build_warmup_reqs( server_args, warmup_resolutions=None, - use_model_sampling_defaults=True, server_based_warmup=True, ) - self.assertEqual((reqs[0].width, reqs[0].height), (1024, 1024)) + self.assertEqual((reqs[0].width, reqs[0].height), (512, 512)) + + def test_server_based_warmup_scales_large_image_default(self): + server_args = MagicMock() + server_args.warmup_steps = 1 + server_args.enable_cfg_parallel = False + server_args.backend = "auto" + + task_type = MagicMock() + task_type.requires_image_input.return_value = False + task_type.accepts_image_input.return_value = False + task_type.is_image_gen.return_value = True + task_type.data_type.return_value = ModelTaskType.T2I.data_type() + server_args.pipeline_config.task_type = task_type + + sampling_defaults = SamplingParams(width=1024, height=1024) + with patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults", + return_value=sampling_defaults, + ): + reqs = build_warmup_reqs( + server_args, + warmup_resolutions=None, + server_based_warmup=True, + ) + + self.assertEqual((reqs[0].width, reqs[0].height), (768, 768)) + + def test_server_based_warmup_uses_diffusers_image_budget(self): + server_args = MagicMock() + server_args.warmup_steps = 1 + server_args.enable_cfg_parallel = False + server_args.backend = "diffusers" + + task_type = MagicMock() + task_type.requires_image_input.return_value = False + task_type.accepts_image_input.return_value = False + task_type.is_image_gen.return_value = True + task_type.data_type.return_value = ModelTaskType.T2I.data_type() + server_args.pipeline_config.task_type = task_type + + sampling_defaults = SamplingParams(width=1024, height=1024) + with patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults", + return_value=sampling_defaults, + ): + reqs = build_warmup_reqs( + server_args, + warmup_resolutions=None, + server_based_warmup=True, + ) + + self.assertEqual((reqs[0].width, reqs[0].height), (512, 512)) def test_server_based_warmup_keeps_video_warmup_lightweight(self): server_args = MagicMock() @@ -270,14 +423,86 @@ class TestWarmupReqCfgParallel(unittest.TestCase): reqs = build_warmup_reqs( server_args, warmup_resolutions=None, - use_model_sampling_defaults=True, server_based_warmup=True, ) - self.assertEqual(reqs[0].num_inference_steps, 1) - self.assertEqual(reqs[0].num_frames, 81) + self.assertEqual(reqs[0].num_inference_steps, 2) + self.assertEqual(reqs[0].num_frames, 17) - def test_server_based_warmup_keeps_lightweight_image_fallback(self): + def test_server_based_warmup_uses_video_supported_resolution_budget(self): + server_args = MagicMock() + server_args.warmup_steps = 1 + server_args.enable_cfg_parallel = False + + task_type = MagicMock() + task_type.requires_image_input.return_value = False + task_type.accepts_image_input.return_value = False + task_type.is_image_gen.return_value = False + task_type.data_type.return_value = ModelTaskType.T2V.data_type() + server_args.pipeline_config.task_type = task_type + + sampling_defaults = SamplingParams( + width=1280, + height=720, + num_frames=81, + num_inference_steps=35, + supported_resolutions=[ + (1280, 720), + (720, 1280), + (832, 480), + (480, 832), + (1024, 1024), + ], + ) + with patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults", + return_value=sampling_defaults, + ): + reqs = build_warmup_reqs( + server_args, + warmup_resolutions=None, + server_based_warmup=True, + ) + + self.assertEqual((reqs[0].width, reqs[0].height), (832, 480)) + self.assertEqual(reqs[0].num_frames, 17) + self.assertEqual(reqs[0].num_inference_steps, 2) + + def test_ltx2_two_stage_warmup_uses_pipeline_alignment(self): + server_args = MagicMock() + server_args.warmup_steps = 1 + server_args.enable_cfg_parallel = False + server_args.pipeline_class_name = "LTX2TwoStageHQPipeline" + + task_type = MagicMock() + task_type.requires_image_input.return_value = False + task_type.accepts_image_input.return_value = False + task_type.is_image_gen.return_value = False + task_type.data_type.return_value = ModelTaskType.T2V.data_type() + server_args.pipeline_config.task_type = task_type + server_args.pipeline_config.vae_scale_factor = 32 + + sampling_defaults = SamplingParams( + width=1920, + height=1088, + num_frames=121, + num_inference_steps=15, + ) + with patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults", + return_value=sampling_defaults, + ): + reqs = build_warmup_reqs( + server_args, + warmup_resolutions=None, + server_based_warmup=True, + ) + + self.assertEqual((reqs[0].width, reqs[0].height), (832, 448)) + self.assertEqual(reqs[0].width % 64, 0) + self.assertEqual(reqs[0].height % 64, 0) + + def test_server_based_warmup_uses_representative_image_fallback(self): server_args = MagicMock() server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False @@ -296,14 +521,13 @@ class TestWarmupReqCfgParallel(unittest.TestCase): reqs = build_warmup_reqs( server_args, warmup_resolutions=None, - use_model_sampling_defaults=True, server_based_warmup=True, ) req = reqs[0] self.assertEqual( (req.width, req.height), - DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION, + SERVER_WARMUP_IMAGE_FALLBACK_RESOLUTION, ) def test_warmup_image_inclusion_policy_all_task_types(self): @@ -346,7 +570,6 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args, warmup_resolutions=None, warmup_input_path="/tmp/warmup.png", - use_model_sampling_defaults=True, server_based_warmup=True, ) @@ -366,7 +589,6 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args, warmup_resolutions=None, warmup_input_path="/tmp/warmup.png", - use_model_sampling_defaults=True, server_based_warmup=True, ) @@ -386,7 +608,6 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args, warmup_resolutions=None, warmup_input_path="/tmp/warmup.png", - use_model_sampling_defaults=True, server_based_warmup=True, )