[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.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_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, build_warmup_reqs,
prepare_warmup_image_path,
should_include_warmup_image, 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.srt.utils.json_response import orjson_response
from sglang.version import __version__ from sglang.version import __version__
@@ -116,7 +116,9 @@ async def _run_server_warmup_after_http_ready(
server_based_warmup=True, server_based_warmup=True,
use_model_sampling_defaults=True, use_model_sampling_defaults=True,
) )
warmup_total = len(warmup_reqs)
for req in warmup_reqs: for req in warmup_reqs:
req.extra["warmup_total"] = warmup_total
response = await async_scheduler_client.forward(req) response = await async_scheduler_client.forward(req)
if response.error is not None: if response.error is not None:
raise RuntimeError(response.error) raise RuntimeError(response.error)
@@ -11,6 +11,7 @@ from enum import Enum
from typing import Any, Iterator, List from typing import Any, Iterator, List
import zmq import zmq
from tqdm.auto import tqdm
from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType
from sglang.multimodal_gen.runtime.disaggregation.scheduler_mixin import ( 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, set_global_server_args,
) )
from sglang.multimodal_gen.runtime.server_warmup import ( from sglang.multimodal_gen.runtime.server_warmup import (
build_warmup_reqs,
get_first_generation_req, get_first_generation_req,
is_server_based_warmup, is_server_based_warmup,
is_warmup_req, is_warmup_req,
prepare_warmup_image_path_sync, prepare_warmup_image_path_sync,
should_include_warmup_image,
should_return_warmup_result, should_return_warmup_result,
) )
from sglang.multimodal_gen.runtime.utils.common import get_zmq_socket 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.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.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__) logger = init_logger(__name__)
@@ -125,6 +128,7 @@ class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin):
self.task_pipes_to_slaves = task_pipes_to_slaves self.task_pipes_to_slaves = task_pipes_to_slaves
self.result_pipes_from_slaves = result_pipes_from_slaves self.result_pipes_from_slaves = result_pipes_from_slaves
self.gpu_id = gpu_id self.gpu_id = gpu_id
self._show_warmup_progress = gpu_id == 0
self._running = True self._running = True
self.request_handlers = { self.request_handlers = {
@@ -160,6 +164,7 @@ class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin):
# warmup progress tracking # warmup progress tracking
self._warmup_total = 0 self._warmup_total = 0
self._warmup_processed = 0 self._warmup_processed = 0
self._warmup_progress_bar: Any | None = None
self._logged_server_ready_after_warmup = False self._logged_server_ready_after_warmup = False
self.prepare_server_warmup_reqs() 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(req) for req in reqs]
return self._dispatch_single_request(reqs[0]) 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( def _log_warmup_result(
self, self,
output_batch: OutputBatch, output_batch: OutputBatch,
@@ -260,23 +334,10 @@ class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin):
return return
server_based_warmup = is_server_based_warmup(req_or_group) 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 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 ( if (
not server_based_warmup not server_based_warmup
and not self._logged_server_ready_after_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!") logger.info("The server is fired up and ready to roll!")
self._logged_server_ready_after_warmup = True self._logged_server_ready_after_warmup = True
else: else:
if self._warmup_total > 0: warmup_desc = self._format_warmup_req(req_or_group)
logger.info( logger.info(f"{warmup_desc} processing failed")
f"Warmup req ({self._warmup_processed}/{self._warmup_total}) processing failed"
)
else:
logger.info("Warmup req processing failed")
def _handle_generation( def _handle_generation(
self, reqs: list[Any], *, allow_dynamic_batching: bool = True self, reqs: list[Any], *, allow_dynamic_batching: bool = True
@@ -302,13 +359,7 @@ class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin):
reqs = self._normalize_generation_reqs(reqs) reqs = self._normalize_generation_reqs(reqs)
warmup_reqs = [req for req in reqs if req.is_warmup] warmup_reqs = [req for req in reqs if req.is_warmup]
if warmup_reqs: if warmup_reqs:
self._warmup_processed += len(warmup_reqs) self._ensure_warmup_progress_bar(warmup_reqs[0])
if self._warmup_total > 0:
logger.info(
f"Processing warmup req... ({self._warmup_processed}/{self._warmup_total})"
)
else:
logger.info("Processing warmup req...")
# Use the head request trace context for scheduler-side dispatch work. # Use the head request trace context for scheduler-side dispatch work.
req = reqs[0] req = reqs[0]
@@ -987,8 +1038,6 @@ class Scheduler(SchedulerPostTrainingMixin, SchedulerDisaggMixin):
): ):
return recv_reqs 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] identity, req_or_group = recv_reqs[0]
req = get_first_generation_req(req_or_group) req = get_first_generation_req(req_or_group)
if req is not None: if req is not None:
@@ -323,6 +323,7 @@ class Req:
self.is_warmup = True self.is_warmup = True
self.save_output = False self.save_output = False
self.suppress_logs = True self.suppress_logs = True
self.metrics.suppress_stage_breakdown = True
self.extra["cache_dit_num_inference_steps"] = self.num_inference_steps self.extra["cache_dit_num_inference_steps"] = self.num_inference_steps
self.num_inference_steps = warmup_steps self.num_inference_steps = warmup_steps
@@ -98,9 +98,11 @@ class PipelineStage(StageDedupMixin, ABC):
total: int | None = None, total: int | None = None,
*, *,
disable: bool = False, disable: bool = False,
batch: Req | None = None,
**kwargs, **kwargs,
) -> tqdm: ) -> tqdm:
is_main_rank = not world_group_is_initialized() or get_world_rank() == 0 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( return tqdm(
iterable=iterable, iterable=iterable,
total=total, total=total,
@@ -398,7 +398,7 @@ class CausalDMDDenoisingStage(DenoisingStage):
def _realtime_causal_progress_bar(self, batch: Req, timesteps: torch.Tensor): def _realtime_causal_progress_bar(self, batch: Req, timesteps: torch.Tensor):
if batch.session is not None: if batch.session is not None:
return nullcontext(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( def _denoise_realtime_causal_chunk(
self, self,
@@ -1025,7 +1025,9 @@ class CausalDMDDenoisingStage(DenoisingStage):
return current_latents return current_latents
# DMD loop in causal blocks # 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: for current_num_frames in block_sizes:
current_latents = latents[ current_latents = latents[
:, :, start_index : start_index + current_num_frames, :, : :, :, start_index : start_index + current_num_frames, :, :
@@ -1375,7 +1375,9 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
), ),
maybe_nvtx_range("denoising_loop", use_nvtx), 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): for step_index, t_host in enumerate(timesteps_cpu):
# Use ``:.4g`` so flow-matching schedulers (e.g. FLUX) that # Use ``:.4g`` so flow-matching schedulers (e.g. FLUX) that
# use non-integer timesteps keep their precision in markers. # use non-integer timesteps keep their precision in markers.
@@ -94,7 +94,7 @@ class DmdDenoisingStage(DenoisingStage):
pos_cond_kwargs = prepared_vars["pos_cond_kwargs"] pos_cond_kwargs = prepared_vars["pos_cond_kwargs"]
denoising_loop_start_time = time.time() 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): for i, t in enumerate(timesteps):
# Skip if interrupted # Skip if interrupted
if hasattr(self, "interrupt") and self.interrupt: if hasattr(self, "interrupt") and self.interrupt:
@@ -616,7 +616,7 @@ class Cosmos3DenoisingStage(PipelineStage):
enumerate(timesteps), enumerate(timesteps),
total=len(timesteps), total=len(timesteps),
desc="Denoising", desc="Denoising",
disable=batch.is_warmup, batch=batch,
) )
for i, t in progress_bar: for i, t in progress_bar:
@@ -467,7 +467,7 @@ class MOVADenoisingStage(PipelineStage):
metrics = getattr(batch, "metrics", None) metrics = getattr(batch, "metrics", None)
perf_dump_path_provided = getattr(batch, "perf_dump_path", None) is not 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): for idx_step in range(total_steps):
with StageProfiler( with StageProfiler(
f"denoising_step_{idx_step}", f"denoising_step_{idx_step}",
@@ -962,7 +962,7 @@ class SanaWMDenoisingStage(DenoisingStage):
assert transformer is not None assert transformer is not None
self.transformer = transformer 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: if cfg_parallel:
latent_model_input = latents latent_model_input = latents
else: else:
@@ -631,7 +631,7 @@ class SanaWMStreamingDenoisingStage(CausalDMDDenoisingStage):
do_cfg, 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( chunk_kv, sink_num = self._accumulate_kv_cache(
kv_cache, kv_cache,
chunk_idx, chunk_idx,
@@ -4,16 +4,9 @@
import asyncio import asyncio
import os import os
import tempfile import tempfile
from copy import copy
from typing import Any from typing import Any
from sglang.multimodal_gen.configs.pipeline_configs.base import ModelTaskType from sglang.multimodal_gen.runtime.entrypoints.openai.utils import save_image_to_path
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.pipelines_core.schedule_batch import Req 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
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger 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==" 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: def get_first_generation_req(req_or_group: Any) -> Req | None:
"""Extract the first req""" """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: async def prepare_warmup_image_path(server_args: ServerArgs) -> str:
if server_args.input_save_path is not None: if server_args.input_save_path is not None:
uploads_dir = server_args.input_save_path 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: def prepare_warmup_image_path_sync(server_args: ServerArgs) -> str:
return asyncio.run(prepare_warmup_image_path(server_args)) 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.stages: Dict[str, float] = {}
self.steps: list[float] = [] self.steps: list[float] = []
self.total_duration_ms: float = 0.0 self.total_duration_ms: float = 0.0
self.suppress_stage_breakdown: bool = False
# memory tracking: {checkpoint_name: MemorySnapshot} # memory tracking: {checkpoint_name: MemorySnapshot}
self.memory_snapshots: Dict[str, MemorySnapshot] = {} self.memory_snapshots: Dict[str, MemorySnapshot] = {}
@@ -62,13 +63,19 @@ class RequestMetrics:
def record_stage(self, stage_name: str, duration_s: float): def record_stage(self, stage_name: str, duration_s: float):
"""Records the duration of a pipeline stage""" """Records the duration of a pipeline stage"""
if self.suppress_stage_breakdown:
return
self.stages[stage_name] = duration_s * 1000 # Store as milliseconds self.stages[stage_name] = duration_s * 1000 # Store as milliseconds
def record_step(self, duration_s: float): def record_step(self, duration_s: float):
"""Records the duration of a denoising step in execution order.""" """Records the duration of a denoising step in execution order."""
if self.suppress_stage_breakdown:
return
self.steps.append(duration_s * 1000) self.steps.append(duration_s * 1000)
def record_memory_snapshot(self, checkpoint_name: str, snapshot: MemorySnapshot): def record_memory_snapshot(self, checkpoint_name: str, snapshot: MemorySnapshot):
if self.suppress_stage_breakdown:
return
self.memory_snapshots[checkpoint_name] = snapshot self.memory_snapshots[checkpoint_name] = snapshot
def to_dict(self) -> Dict[str, Any]: 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. """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 - Scheduler.prepare_server_warmup_reqs synthesizes warmup Reqs that
actually enable classifier-free guidance when cfg-parallel is on. actually enable classifier-free guidance when cfg-parallel is on.
- InputValidationStage.forward rejects non-CFG requests when the server - InputValidationStage.forward rejects non-CFG requests when the server
has cfg-parallel on. has cfg-parallel on.
- Server-based warmup can opt into model-default negative prompts so warmup - Server-based warmup can opt into model-default negative prompts so warmup
populates the negative text embedding cache. 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. 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 ( from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import (
InputValidationStage, InputValidationStage,
) )
from sglang.multimodal_gen.runtime.server_warmup import ( from sglang.multimodal_gen.runtime.warmup_request_builder import (
DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION, DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION,
DEFAULT_PLACEHOLDER_PROMPT, DEFAULT_PLACEHOLDER_PROMPT,
build_warmup_reqs, build_warmup_reqs,
@@ -73,6 +74,16 @@ def _make_input_validation_stage() -> InputValidationStage:
return 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: def _make_validation_server_args(enable_cfg_parallel: bool) -> MagicMock:
sa = MagicMock() sa = MagicMock()
sa.enable_cfg_parallel = enable_cfg_parallel sa.enable_cfg_parallel = enable_cfg_parallel
@@ -106,6 +117,36 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
self.assertIs(req.do_classifier_free_guidance, False) self.assertIs(req.do_classifier_free_guidance, False)
self.assertNotEqual(req.negative_prompt, DEFAULT_PLACEHOLDER_PROMPT) 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): def test_server_based_warmup_uses_model_default_negative_prompt(self):
server_args = MagicMock() server_args = MagicMock()
server_args.warmup_steps = 1 server_args.warmup_steps = 1
@@ -124,7 +165,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
num_inference_steps=20, num_inference_steps=20,
) )
with patch( 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, return_value=sampling_defaults,
): ):
reqs = build_warmup_reqs( reqs = build_warmup_reqs(
@@ -138,6 +179,9 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
self.assertEqual(len(reqs), 1) self.assertEqual(len(reqs), 1)
req = reqs[0] req = reqs[0]
self.assertTrue(req.is_warmup) 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.assertEqual(req.negative_prompt, "model default negative")
self.assertIs(req.do_classifier_free_guidance, True) self.assertIs(req.do_classifier_free_guidance, True)
self.assertTrue(req.extra["return_warmup_result"]) self.assertTrue(req.extra["return_warmup_result"])
@@ -157,7 +201,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
sampling_defaults = SamplingParams(width=640, height=640) sampling_defaults = SamplingParams(width=640, height=640)
with patch( 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, return_value=sampling_defaults,
): ):
reqs = build_warmup_reqs( reqs = build_warmup_reqs(
@@ -171,6 +215,68 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
self.assertEqual(req.width, 640) self.assertEqual(req.width, 640)
self.assertEqual(req.height, 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): def test_server_based_warmup_keeps_lightweight_image_fallback(self):
server_args = MagicMock() server_args = MagicMock()
server_args.warmup_steps = 1 server_args.warmup_steps = 1
@@ -184,7 +290,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.pipeline_config.task_type = task_type server_args.pipeline_config.task_type = task_type
with patch( 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(), return_value=SamplingParams(),
): ):
reqs = build_warmup_reqs( reqs = build_warmup_reqs(
@@ -233,7 +339,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.pipeline_config.task_type = ModelTaskType.TI2I server_args.pipeline_config.task_type = ModelTaskType.TI2I
with patch( 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), return_value=SamplingParams(width=512, height=512),
): ):
reqs = build_warmup_reqs( reqs = build_warmup_reqs(
@@ -253,7 +359,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.pipeline_config.task_type = ModelTaskType.I2I server_args.pipeline_config.task_type = ModelTaskType.I2I
with patch( 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), return_value=SamplingParams(width=512, height=512),
): ):
reqs = build_warmup_reqs( reqs = build_warmup_reqs(
@@ -273,7 +379,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.pipeline_config.task_type = ModelTaskType.TI2V server_args.pipeline_config.task_type = ModelTaskType.TI2V
with patch( 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), return_value=SamplingParams(width=512, height=512),
): ):
reqs = build_warmup_reqs( reqs = build_warmup_reqs(