[diffusion] chore: improve server warmup coverage (#28127)

This commit is contained in:
Mick
2026-06-14 13:35:42 +08:00
committed by GitHub
parent 5fb4e2d02e
commit 31ac743484
11 changed files with 787 additions and 417 deletions
@@ -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():
@@ -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:
@@ -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:[<media-type>];base64,<data>`"
)
# split `data:[<media-type>][;base64],<data>` 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,
@@ -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
@@ -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
@@ -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",
@@ -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
@@ -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")
@@ -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:[<media-type>];base64,<data>`"
)
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
@@ -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)
@@ -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,
)