[diffusion] warmup: improve diffusion server warmup (#28119)

This commit is contained in:
Mick
2026-06-13 13:04:10 +08:00
committed by GitHub
parent eb18416f9f
commit 8becb37519
15 changed files with 424 additions and 205 deletions
@@ -35,12 +35,12 @@ 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 (
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,
prepare_warmup_image_path,
should_include_warmup_image,
)
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__
@@ -116,7 +116,9 @@ async def _run_server_warmup_after_http_ready(
server_based_warmup=True,
use_model_sampling_defaults=True,
)
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)
@@ -11,6 +11,7 @@ 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 (
@@ -55,18 +56,20 @@ from sglang.multimodal_gen.runtime.server_args import (
set_global_server_args,
)
from sglang.multimodal_gen.runtime.server_warmup import (
build_warmup_reqs,
get_first_generation_req,
is_server_based_warmup,
is_warmup_req,
prepare_warmup_image_path_sync,
should_include_warmup_image,
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 GREEN, RESET, init_logger
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__)
@@ -125,6 +128,7 @@ class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin):
self.task_pipes_to_slaves = task_pipes_to_slaves
self.result_pipes_from_slaves = result_pipes_from_slaves
self.gpu_id = gpu_id
self._show_warmup_progress = gpu_id == 0
self._running = True
self.request_handlers = {
@@ -160,6 +164,7 @@ class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin):
# 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()
@@ -250,6 +255,75 @@ 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,
@@ -260,23 +334,10 @@ class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin):
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:
total_duration_s = (
output_batch.metrics.total_duration_s
if output_batch.metrics is not None
else 0.0
)
if self._warmup_total > 0:
logger.info(
f"Warmup req ({self._warmup_processed}/{self._warmup_total}) processed in {GREEN}%.2f{RESET} seconds",
total_duration_s,
)
else:
logger.info(
f"Warmup req processed in {GREEN}%.2f{RESET} seconds",
total_duration_s,
)
if (
not server_based_warmup
and not self._logged_server_ready_after_warmup
@@ -288,12 +349,8 @@ class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin):
logger.info("The server is fired up and ready to roll!")
self._logged_server_ready_after_warmup = True
else:
if self._warmup_total > 0:
logger.info(
f"Warmup req ({self._warmup_processed}/{self._warmup_total}) processing failed"
)
else:
logger.info("Warmup req processing failed")
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
@@ -302,13 +359,7 @@ class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin):
reqs = self._normalize_generation_reqs(reqs)
warmup_reqs = [req for req in reqs if req.is_warmup]
if warmup_reqs:
self._warmup_processed += len(warmup_reqs)
if self._warmup_total > 0:
logger.info(
f"Processing warmup req... ({self._warmup_processed}/{self._warmup_total})"
)
else:
logger.info("Processing warmup req...")
self._ensure_warmup_progress_bar(warmup_reqs[0])
# Use the head request trace context for scheduler-side dispatch work.
req = reqs[0]
@@ -987,8 +1038,6 @@ class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin):
):
return recv_reqs
# handle server req-based warmup by inserting an identical req to the beginning of the waiting queue
# only the very first req through server's lifetime will be warmed up
identity, req_or_group = recv_reqs[0]
req = get_first_generation_req(req_or_group)
if req is not None:
@@ -323,6 +323,7 @@ class Req:
self.is_warmup = True
self.save_output = False
self.suppress_logs = True
self.metrics.suppress_stage_breakdown = True
self.extra["cache_dit_num_inference_steps"] = self.num_inference_steps
self.num_inference_steps = warmup_steps
@@ -98,9 +98,11 @@ class PipelineStage(StageDedupMixin, ABC):
total: int | None = None,
*,
disable: bool = False,
batch: Req | None = None,
**kwargs,
) -> tqdm:
is_main_rank = not world_group_is_initialized() or get_world_rank() == 0
disable = disable or (batch is not None and batch.is_warmup)
return tqdm(
iterable=iterable,
total=total,
@@ -398,7 +398,7 @@ class CausalDMDDenoisingStage(DenoisingStage):
def _realtime_causal_progress_bar(self, batch: Req, timesteps: torch.Tensor):
if batch.session is not None:
return nullcontext(None)
return self.progress_bar(total=len(timesteps))
return self.progress_bar(total=len(timesteps), batch=batch)
def _denoise_realtime_causal_chunk(
self,
@@ -1025,7 +1025,9 @@ class CausalDMDDenoisingStage(DenoisingStage):
return current_latents
# DMD loop in causal blocks
with self.progress_bar(total=len(block_sizes) * len(timesteps)) as progress_bar:
with self.progress_bar(
total=len(block_sizes) * len(timesteps), batch=batch
) as progress_bar:
for current_num_frames in block_sizes:
current_latents = latents[
:, :, start_index : start_index + current_num_frames, :, :
@@ -1375,7 +1375,9 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
),
maybe_nvtx_range("denoising_loop", use_nvtx),
):
with self.progress_bar(total=ctx.num_inference_steps) as progress_bar:
with self.progress_bar(
total=ctx.num_inference_steps, batch=batch
) as progress_bar:
for step_index, t_host in enumerate(timesteps_cpu):
# Use ``:.4g`` so flow-matching schedulers (e.g. FLUX) that
# use non-integer timesteps keep their precision in markers.
@@ -94,7 +94,7 @@ class DmdDenoisingStage(DenoisingStage):
pos_cond_kwargs = prepared_vars["pos_cond_kwargs"]
denoising_loop_start_time = time.time()
with self.progress_bar(total=len(timesteps)) as progress_bar:
with self.progress_bar(total=len(timesteps), batch=batch) as progress_bar:
for i, t in enumerate(timesteps):
# Skip if interrupted
if hasattr(self, "interrupt") and self.interrupt:
@@ -616,7 +616,7 @@ class Cosmos3DenoisingStage(PipelineStage):
enumerate(timesteps),
total=len(timesteps),
desc="Denoising",
disable=batch.is_warmup,
batch=batch,
)
for i, t in progress_bar:
@@ -467,7 +467,7 @@ class MOVADenoisingStage(PipelineStage):
metrics = getattr(batch, "metrics", None)
perf_dump_path_provided = getattr(batch, "perf_dump_path", None) is not None
with self.progress_bar(total=total_steps) as progress_bar:
with self.progress_bar(total=total_steps, batch=batch) as progress_bar:
for idx_step in range(total_steps):
with StageProfiler(
f"denoising_step_{idx_step}",
@@ -962,7 +962,7 @@ class SanaWMDenoisingStage(DenoisingStage):
assert transformer is not None
self.transformer = transformer
for step_idx, t in enumerate(self.progress_bar(timesteps)):
for step_idx, t in enumerate(self.progress_bar(timesteps, batch=batch)):
if cfg_parallel:
latent_model_input = latents
else:
@@ -631,7 +631,7 @@ class SanaWMStreamingDenoisingStage(CausalDMDDenoisingStage):
do_cfg,
)
for chunk_idx in self.progress_bar(range(num_chunks)):
for chunk_idx in self.progress_bar(range(num_chunks), batch=batch):
chunk_kv, sink_num = self._accumulate_kv_cache(
kv_cache,
chunk_idx,
@@ -4,16 +4,9 @@
import asyncio
import os
import tempfile
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.registry import get_pipeline_config_classes
from sglang.multimodal_gen.runtime.entrypoints.openai.utils import (
_parse_size,
save_image_to_path,
)
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 sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -22,9 +15,6 @@ logger = init_logger(__name__)
MINIMUM_PICTURE_BASE64_FOR_WARMUP = "data:image/jpg;base64,iVBORw0KGgoAAAANSUhEUgAAACAAAAAgCAYAAABzenr0AAAACXBIWXMAAA7EAAAOxAGVKw4bAAAAbUlEQVRYhe3VsQ2AMAxE0Y/lIgNQULD/OqyCMgCihCKSG4yRuKuiNH6JLsoEbMACOGBcua9HOR7Y6w6swBwMy0qLTpkeI77qdEBpBFAHBBDAGH8WrwJKI4AAegUCfAKgEgpQDvh3CR3oQCuav58qlAw73kKCSgAAAABJRU5ErkJggg=="
DEFAULT_PLACEHOLDER_PROMPT = "warmup"
DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION = (64, 64)
def get_first_generation_req(req_or_group: Any) -> Req | None:
"""Extract the first req"""
@@ -60,25 +50,6 @@ def should_return_warmup_result(req_or_group: Any) -> bool:
)
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()
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()
async 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
@@ -94,126 +65,3 @@ async def prepare_warmup_image_path(server_args: ServerArgs) -> str:
def prepare_warmup_image_path_sync(server_args: ServerArgs) -> str:
return asyncio.run(prepare_warmup_image_path(server_args))
def _resolve_default_warmup_resolution(
server_args: ServerArgs,
sampling_defaults: SamplingParams,
) -> tuple[int, int]:
supported_resolutions = sampling_defaults.supported_resolutions
if supported_resolutions:
return min(supported_resolutions, key=lambda size: size[0] * size[1])
width = sampling_defaults.width
height = sampling_defaults.height
if width is not None and height is not None:
return width, height
if server_args.pipeline_config.task_type.is_image_gen():
return DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION
return (
width or DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION[0],
height or DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION[1],
)
def _effective_cfg_scale(sampling_defaults: SamplingParams) -> float | None:
if sampling_defaults.true_cfg_scale is not None:
return sampling_defaults.true_cfg_scale
return sampling_defaults.guidance_scale
def should_include_warmup_image(
server_args: ServerArgs, server_based_warmup: bool
) -> bool:
task_type = server_args.pipeline_config.task_type
if not task_type.accepts_image_input():
return False
if task_type.requires_image_input():
return True
if server_based_warmup:
return task_type in (ModelTaskType.TI2I, ModelTaskType.TI2V)
return True
def build_warmup_reqs(
server_args: ServerArgs,
*,
warmup_resolutions: list[str] | None,
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()
if warmup_resolutions is None:
width, height = _resolve_default_warmup_resolution(
server_args, sampling_defaults
)
resolutions: list[tuple[int, int]] = [(width, height)]
else:
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
)
warmup_reqs = []
include_warmup_image = should_include_warmup_image(server_args, server_based_warmup)
for width, height in resolutions:
req_kwargs = dict(
data_type=task_type.data_type(),
width=width,
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,
)
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
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
):
req_kwargs["do_classifier_free_guidance"] = True
req = Req(**req_kwargs)
req.set_as_warmup(server_args.warmup_steps)
if return_warmup_result:
req.extra["return_warmup_result"] = True
if server_based_warmup:
req.extra["server_based_warmup"] = True
warmup_reqs.append(req)
return warmup_reqs
@@ -53,6 +53,7 @@ class RequestMetrics:
self.stages: Dict[str, float] = {}
self.steps: list[float] = []
self.total_duration_ms: float = 0.0
self.suppress_stage_breakdown: bool = False
# memory tracking: {checkpoint_name: MemorySnapshot}
self.memory_snapshots: Dict[str, MemorySnapshot] = {}
@@ -62,13 +63,19 @@ class RequestMetrics:
def record_stage(self, stage_name: str, duration_s: float):
"""Records the duration of a pipeline stage"""
if self.suppress_stage_breakdown:
return
self.stages[stage_name] = duration_s * 1000 # Store as milliseconds
def record_step(self, duration_s: float):
"""Records the duration of a denoising step in execution order."""
if self.suppress_stage_breakdown:
return
self.steps.append(duration_s * 1000)
def record_memory_snapshot(self, checkpoint_name: str, snapshot: MemorySnapshot):
if self.suppress_stage_breakdown:
return
self.memory_snapshots[checkpoint_name] = snapshot
def to_dict(self) -> Dict[str, Any]:
@@ -0,0 +1,200 @@
# 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.
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.
"""
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.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.utils.logging_utils import init_logger
logger = init_logger(__name__)
DEFAULT_PLACEHOLDER_PROMPT = "warmup"
DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION = (64, 64)
SERVER_WARMUP_IMAGE_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()
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()
def _resolve_default_warmup_resolution(
server_args: ServerArgs,
sampling_defaults: SamplingParams,
) -> tuple[int, int]:
width = sampling_defaults.width
height = sampling_defaults.height
if width is not None and height is not None:
return width, height
supported_resolutions = sampling_defaults.supported_resolutions
if supported_resolutions:
return min(supported_resolutions, key=lambda size: size[0] * size[1])
if server_args.pipeline_config.task_type.is_image_gen():
return DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION
return (
width or DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION[0],
height or DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION[1],
)
def _effective_cfg_scale(sampling_defaults: SamplingParams) -> float | None:
if sampling_defaults.true_cfg_scale is not None:
return sampling_defaults.true_cfg_scale
return sampling_defaults.guidance_scale
def _resolve_warmup_steps(
server_args: ServerArgs,
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:
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
return min(default_steps, max(warmup_steps, SERVER_WARMUP_IMAGE_STEPS))
def should_include_warmup_image(
server_args: ServerArgs, server_based_warmup: bool
) -> bool:
task_type = server_args.pipeline_config.task_type
if not task_type.accepts_image_input():
return False
if task_type.requires_image_input():
return True
if server_based_warmup:
return task_type in (ModelTaskType.TI2I, ModelTaskType.TI2V)
return True
def build_warmup_reqs(
server_args: ServerArgs,
*,
warmup_resolutions: list[str] | None,
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()
if warmup_resolutions is None:
width, height = _resolve_default_warmup_resolution(
server_args, sampling_defaults
)
resolutions: list[tuple[int, int]] = [(width, height)]
else:
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
)
warmup_steps = _resolve_warmup_steps(
server_args,
sampling_defaults,
server_based_warmup=server_based_warmup,
use_model_sampling_defaults=use_model_sampling_defaults,
)
warmup_reqs = []
include_warmup_image = should_include_warmup_image(server_args, server_based_warmup)
for width, height in resolutions:
req_kwargs = dict(
data_type=task_type.data_type(),
width=width,
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,
)
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
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
):
req_kwargs["do_classifier_free_guidance"] = True
req = Req(**req_kwargs)
req.set_as_warmup(warmup_steps)
if return_warmup_result:
req.extra["return_warmup_result"] = True
if server_based_warmup:
req.extra["server_based_warmup"] = True
warmup_reqs.append(req)
return warmup_reqs
@@ -1,12 +1,13 @@
"""Unit tests for the --enable-cfg-parallel warmup fix and guard.
Covers three code paths introduced alongside this file:
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.
- 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.
All tests are CPU-only; no model loading, no distributed init.
"""
@@ -31,7 +32,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import (
from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import (
InputValidationStage,
)
from sglang.multimodal_gen.runtime.server_warmup import (
from sglang.multimodal_gen.runtime.warmup_request_builder import (
DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION,
DEFAULT_PLACEHOLDER_PROMPT,
build_warmup_reqs,
@@ -73,6 +74,16 @@ def _make_input_validation_stage() -> InputValidationStage:
return InputValidationStage()
def _make_generation_req() -> Req:
return Req(
data_type=ModelTaskType.T2I.data_type(),
prompt="prompt",
width=512,
height=512,
num_inference_steps=20,
)
def _make_validation_server_args(enable_cfg_parallel: bool) -> MagicMock:
sa = MagicMock()
sa.enable_cfg_parallel = enable_cfg_parallel
@@ -106,6 +117,36 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
self.assertIs(req.do_classifier_free_guidance, False)
self.assertNotEqual(req.negative_prompt, DEFAULT_PLACEHOLDER_PROMPT)
def test_req_based_warmup_remains_explicit_legacy_entry(self):
scheduler = _make_bare_scheduler(enable_cfg_parallel=False)
scheduler.server_args.warmup_resolutions = None
scheduler.server_args.server_warmup = False
req = _make_generation_req()
recv_reqs = [(b"0", req)]
processed = scheduler.process_received_reqs_with_req_based_warmup(recv_reqs)
self.assertEqual(len(processed), 2)
self.assertIs(processed[1][1], req)
self.assertIsNot(processed[0][1], req)
self.assertTrue(processed[0][1].is_warmup)
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)
def test_req_based_warmup_skips_default_server_warmup_path(self):
scheduler = _make_bare_scheduler(enable_cfg_parallel=False)
scheduler.server_args.warmup_resolutions = None
scheduler.server_args.server_warmup = True
recv_reqs = [(b"0", _make_generation_req())]
processed = scheduler.process_received_reqs_with_req_based_warmup(recv_reqs)
self.assertIs(processed, recv_reqs)
self.assertEqual(len(processed), 1)
self.assertFalse(scheduler.warmed_up)
def test_server_based_warmup_uses_model_default_negative_prompt(self):
server_args = MagicMock()
server_args.warmup_steps = 1
@@ -124,7 +165,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
num_inference_steps=20,
)
with patch(
"sglang.multimodal_gen.runtime.server_warmup.get_model_sampling_defaults",
"sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults",
return_value=sampling_defaults,
):
reqs = build_warmup_reqs(
@@ -138,6 +179,9 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
self.assertEqual(len(reqs), 1)
req = reqs[0]
self.assertTrue(req.is_warmup)
self.assertTrue(req.metrics.suppress_stage_breakdown)
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.assertIs(req.do_classifier_free_guidance, True)
self.assertTrue(req.extra["return_warmup_result"])
@@ -157,7 +201,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
sampling_defaults = SamplingParams(width=640, height=640)
with patch(
"sglang.multimodal_gen.runtime.server_warmup.get_model_sampling_defaults",
"sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults",
return_value=sampling_defaults,
):
reqs = build_warmup_reqs(
@@ -171,6 +215,68 @@ 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):
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 = 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,
supported_resolutions=[(512, 512), (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,
use_model_sampling_defaults=True,
server_based_warmup=True,
)
self.assertEqual((reqs[0].width, reqs[0].height), (1024, 1024))
def test_server_based_warmup_keeps_video_warmup_lightweight(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=832,
height=480,
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=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)
def test_server_based_warmup_keeps_lightweight_image_fallback(self):
server_args = MagicMock()
server_args.warmup_steps = 1
@@ -184,7 +290,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.pipeline_config.task_type = task_type
with patch(
"sglang.multimodal_gen.runtime.server_warmup.get_model_sampling_defaults",
"sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults",
return_value=SamplingParams(),
):
reqs = build_warmup_reqs(
@@ -233,7 +339,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.pipeline_config.task_type = ModelTaskType.TI2I
with patch(
"sglang.multimodal_gen.runtime.server_warmup.get_model_sampling_defaults",
"sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults",
return_value=SamplingParams(width=512, height=512),
):
reqs = build_warmup_reqs(
@@ -253,7 +359,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.pipeline_config.task_type = ModelTaskType.I2I
with patch(
"sglang.multimodal_gen.runtime.server_warmup.get_model_sampling_defaults",
"sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults",
return_value=SamplingParams(width=512, height=512),
):
reqs = build_warmup_reqs(
@@ -273,7 +379,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.pipeline_config.task_type = ModelTaskType.TI2V
with patch(
"sglang.multimodal_gen.runtime.server_warmup.get_model_sampling_defaults",
"sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults",
return_value=SamplingParams(width=512, height=512),
):
reqs = build_warmup_reqs(