[diffusion] warmup: improve diffusion server warmup (#28119)
This commit is contained in:
@@ -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:
|
||||||
|
|||||||
+1
-1
@@ -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:
|
||||||
|
|||||||
+1
-1
@@ -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}",
|
||||||
|
|||||||
+1
-1
@@ -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:
|
||||||
|
|||||||
+1
-1
@@ -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(
|
||||||
|
|||||||
Reference in New Issue
Block a user