[diffusion] chore: improve server warmup coverage (#28127)
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user