From 03962d4238ab248606b2fd561b251425fbc6fd19 Mon Sep 17 00:00:00 2001 From: Mick Date: Sat, 4 Jul 2026 15:25:26 +0800 Subject: [PATCH] [diffusion] feat: enable compile warmup for vae decode (#29306) --- .../configs/pipeline_configs/zimage.py | 21 +++- .../runtime/managers/scheduler.py | 14 ++- .../runtime/models/dits/zimage.py | 28 ++++- .../runtime/pipelines_core/stages/decoding.py | 35 +++++- .../pipelines_core/stages/denoising.py | 42 ++++---- .../multimodal_gen/runtime/server_args.py | 18 +++- .../multimodal_gen/runtime/server_warmup.py | 5 +- .../runtime/utils/torch_compile.py | 102 ++++++++++++++++++ .../runtime/warmup_request_builder.py | 38 +++++-- .../test/unit/test_cfg_parallel_warmup.py | 97 +++++++++++++++++ .../unit/test_decoding_stage_parallelism.py | 31 ++++++ .../test/unit/test_server_args.py | 38 +++++++ .../test/unit/test_zimage_pipeline_config.py | 2 + 13 files changed, 428 insertions(+), 43 deletions(-) create mode 100644 python/sglang/multimodal_gen/runtime/utils/torch_compile.py diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/zimage.py b/python/sglang/multimodal_gen/configs/pipeline_configs/zimage.py index 441c78abf..5cc5b5fd9 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/zimage.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/zimage.py @@ -202,6 +202,12 @@ class ZImagePipelineConfig(ZImageRolloutPipelineMixin, ImagePipelineConfig): plan = self._build_zimage_sp_plan(batch) return plan + def _pad_text_embed_for_dit(self, embed: torch.Tensor) -> torch.Tensor: + target_len = self._ceil_to_multiple(embed.shape[0], self.SEQ_LEN_MULTIPLE) + if target_len == embed.shape[0]: + return embed + return torch.cat([embed, embed[-1:].repeat(target_len - embed.shape[0], 1)]) + def _split_text_embeds_for_dit(self, batch, *, negative: bool = False): """Return per-request text tensors, trimming padded batched embeddings.""" embeds = batch.negative_prompt_embeds if negative else batch.prompt_embeds @@ -217,7 +223,7 @@ class ZImagePipelineConfig(ZImageRolloutPipelineMixin, ImagePipelineConfig): return embeds if embeds.ndim == 2: - return [embeds] + return [self._pad_text_embed_for_dit(embeds)] if embeds.ndim != 3: raise ValueError( @@ -231,7 +237,8 @@ class ZImagePipelineConfig(ZImageRolloutPipelineMixin, ImagePipelineConfig): expected_batch_size=int(embeds.shape[0]), ) return [ - embeds[idx, :seq_len].contiguous() for idx, seq_len in enumerate(seq_lens) + self._pad_text_embed_for_dit(embeds[idx, :seq_len].contiguous()) + for idx, seq_len in enumerate(seq_lens) ] def _caption_rope_length(self, prompt_embeds, batch, *, negative: bool = False): @@ -465,6 +472,11 @@ class ZImagePipelineConfig(ZImageRolloutPipelineMixin, ImagePipelineConfig): if get_sp_world_size() > 1 else None ), + "caption_valid_lens": torch.tensor( + self.require_text_seq_lens(batch, 0), + device=device, + dtype=torch.long, + ), } def prepare_neg_cond_kwargs(self, batch, device, rotary_emb, dtype): @@ -489,4 +501,9 @@ class ZImagePipelineConfig(ZImageRolloutPipelineMixin, ImagePipelineConfig): if get_sp_world_size() > 1 else None ), + "caption_valid_lens": torch.tensor( + self.require_text_seq_lens(batch, 0, negative=use_negative_embeds), + device=device, + dtype=torch.long, + ), } diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index 11cce5294..3562ecd1a 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -57,6 +57,7 @@ from sglang.multimodal_gen.runtime.server_args import ( ) from sglang.multimodal_gen.runtime.server_warmup import ( SchedulerWarmupMixin, + get_first_generation_req, is_warmup_req, should_return_warmup_result, ) @@ -571,6 +572,12 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag ) -> List[OutputBatch]: return [OutputBatch(error=error_msg) for _ in reqs] + def _should_return_lightweight_warmup_result(self, processed_req: Any) -> bool: + req = get_first_generation_req(processed_req) + return (req is not None and bool(req.extra.get("server_internal_prewarm"))) or ( + is_warmup_req(processed_req) and should_return_warmup_result(processed_req) + ) + def return_result( self, output_batch: OutputBatch, @@ -1048,8 +1055,11 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag is_warmup = is_warmup_req(processed_req) self._log_warmup_result(output_batch, processed_req, is_warmup) - if is_warmup and should_return_warmup_result(processed_req): - # only keep the necessary lightweight payloads + should_return_lightweight_warmup_result = ( + self._should_return_lightweight_warmup_result(processed_req) + ) + if should_return_lightweight_warmup_result: + # internal prewarm is a real-path request; reply but drop payloads output_batch.drop_payload_for_warmup() self.return_result( output_batch, identity, should_not_return=False diff --git a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py index 71adf9962..eda6ea3e0 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/zimage.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/zimage.py @@ -817,6 +817,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): patch_size: int, f_patch_size: int, image_seq_len_target: int | None = None, + caption_valid_lens: torch.Tensor | None = None, ): """Patchify images and pad image/caption tokens to batch targets. @@ -846,7 +847,12 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): for cap_feat in all_cap_feats ) - for cap_feat in all_cap_feats: + if caption_valid_lens is not None: + caption_valid_lens = caption_valid_lens.to( + device=all_cap_feats[0].device, dtype=torch.long + ) + + for idx, cap_feat in enumerate(all_cap_feats): cap_ori_len = cap_feat.size(0) cap_padding_len = cap_seq_len_target - cap_ori_len cap_padded_feat = torch.cat( @@ -854,7 +860,10 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): dim=0, ) all_cap_feats_out.append(cap_padded_feat) - all_cap_valid_lens.append(cap_ori_len) + if caption_valid_lens is None: + all_cap_valid_lens.append(cap_ori_len) + else: + all_cap_valid_lens.append(caption_valid_lens[idx]) target_image_seq_len = image_seq_len_target or 0 for image in all_image: @@ -885,12 +894,15 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): all_image_size.append(image_size) all_image_valid_lens.append(image_ori_len) + cap_valid_lens_out = ( + caption_valid_lens if caption_valid_lens is not None else all_cap_valid_lens + ) return ( torch.stack(all_image_out, dim=0), torch.stack(all_cap_feats_out, dim=0), all_image_size, all_image_valid_lens, - all_cap_valid_lens, + cap_valid_lens_out, ) @staticmethod @@ -923,12 +935,16 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): @staticmethod def _replace_padding_with_token( tensor: torch.Tensor, - valid_lens: list[int], + valid_lens: list[int] | torch.Tensor, pad_token: torch.Tensor, ) -> torch.Tensor: """Replace padded token rows after each valid sequence length.""" positions = torch.arange(tensor.shape[1], device=tensor.device).unsqueeze(0) - lengths = torch.tensor(valid_lens, device=tensor.device).unsqueeze(1) + if torch.is_tensor(valid_lens): + lengths = valid_lens.to(device=tensor.device, dtype=torch.long) + else: + lengths = torch.tensor(valid_lens, device=tensor.device) + lengths = lengths.unsqueeze(1) pad_mask = positions >= lengths if pad_mask.any(): tensor = tensor.clone() @@ -945,6 +961,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): f_patch_size=1, freqs_cis=None, image_seq_len_target: int | None = None, + caption_valid_lens: torch.Tensor | None = None, **kwargs, ): assert patch_size in self.all_patch_size @@ -968,6 +985,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): patch_size, f_patch_size, image_seq_len_target=image_seq_len_target, + caption_valid_lens=caption_valid_lens, ) x, _ = self.all_x_embedder[f"{patch_size}-{f_patch_size}"](x) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py index 1dd0a1623..5474bd3ff 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py @@ -8,6 +8,7 @@ Decoding stage for diffusion pipelines. import weakref import torch +import torch.nn as nn from sglang.multimodal_gen.runtime.distributed import ( get_decode_parallel_world_size, @@ -35,6 +36,11 @@ from sglang.multimodal_gen.runtime.utils.precision import ( resolve_precision, temporary_module_dtype, ) +from sglang.multimodal_gen.runtime.utils.torch_compile import ( + ActiveTargetCompiledCallable, + build_torch_compile_kwargs, + resolve_torch_compile_mode, +) logger = init_logger(__name__) @@ -105,6 +111,7 @@ class DecodingStage(PipelineStage): self.vae: ParallelTiledVAE = vae self.pipeline = weakref.ref(pipeline) if pipeline else None self.component_name = component_name + self._compiled_vae_decode = ActiveTargetCompiledCallable() def component_uses( self, server_args: ServerArgs, stage_name: str | None = None @@ -155,6 +162,32 @@ class DecodingStage(PipelineStage): def scale_and_shift(self, latents: torch.Tensor, server_args): return scale_and_shift_latents(latents, server_args, self.vae) + def _get_vae_decode_fn(self, vae, server_args: ServerArgs): + if not server_args.enable_torch_compile or not isinstance(vae, nn.Module): + return vae.decode + + will_compile = ( + self._compiled_vae_decode.target_id != id(vae) + or self._compiled_vae_decode.compiled_module is None + ) + if current_platform.is_npu(): + compile_kwargs = build_torch_compile_kwargs(mode=None) + if will_compile: + logger.info("Compiling VAE decode with torchair backend on NPU") + else: + mode = resolve_torch_compile_mode( + "SGLANG_VAE_TORCH_COMPILE_MODE", + "SGLANG_TORCH_COMPILE_MODE", + default="default", + ) + compile_kwargs = build_torch_compile_kwargs(mode=mode) + if will_compile: + logger.info("Compiling VAE decode with mode: %s", mode) + + return self._compiled_vae_decode.get_or_compile( + vae, vae.decode, compile_kwargs=compile_kwargs + ) + @torch.no_grad() def decode( self, @@ -209,7 +242,7 @@ class DecodingStage(PipelineStage): with temporary_module_dtype( self.vae, vae_dtype, enabled=should_cast_vae ) as vae: - decode_output = vae.decode(latents) + decode_output = self._get_vae_decode_fn(vae, server_args)(latents) image = _ensure_tensor_decode_output(decode_output) # De-normalize image to [0, 1] range diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index 56ddb8b54..601769eef 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -7,7 +7,6 @@ Denoising stage for diffusion pipelines. import inspect import math -import os import time import weakref from collections.abc import Callable @@ -114,8 +113,13 @@ from sglang.multimodal_gen.runtime.utils.precision import ( resolve_precision, ) from sglang.multimodal_gen.runtime.utils.profiler import SGLDiffusionProfiler +from sglang.multimodal_gen.runtime.utils.torch_compile import ( + CompiledModuleRegistry, + build_torch_compile_kwargs, + maybe_enable_inductor_compute_comm_overlap, + resolve_torch_compile_mode, +) from sglang.multimodal_gen.utils import dict_to_3d_list -from sglang.srt.utils.common import get_compiler_backend logger = init_logger(__name__) @@ -209,7 +213,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): # cache-dit state (for delayed mounting and idempotent control) self._cache_dit_enabled = False self._cached_num_steps = None - self._torch_compiled_module_ids: set[int] = set() + self._torch_compile_registry = CompiledModuleRegistry() hidden_size = self.server_args.pipeline_config.dit_config.hidden_size num_attention_heads = ( @@ -373,31 +377,21 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): if envs.SGLANG_CACHE_DIT_ENABLED and not self._cache_dit_enabled: logger.debug("Deferring torch.compile until cache-dit is enabled") return - module_id = id(module) - if module_id in self._torch_compiled_module_ids: + if self._torch_compile_registry.is_compiled(module): return - compile_kwargs: dict[str, Any] = {"fullgraph": False, "dynamic": None} - if current_platform.is_npu(): - backend = get_compiler_backend() - compile_kwargs["backend"] = backend - compile_kwargs["dynamic"] = False + compile_kwargs = build_torch_compile_kwargs(mode=None) logger.info("Compiling transformer with torchair backend on NPU") else: - try: - import torch._inductor.config as _inductor_cfg - - _inductor_cfg.reorder_for_compute_comm_overlap = True - except ImportError: - pass + maybe_enable_inductor_compute_comm_overlap() dit_config = getattr(self.server_args.pipeline_config, "dit_config", None) - mode = os.environ.get("SGLANG_TORCH_COMPILE_MODE") or getattr( - dit_config, - "torch_compile_mode", - "max-autotune-no-cudagraphs", + mode = resolve_torch_compile_mode( + "SGLANG_TORCH_COMPILE_MODE", + config=dit_config, + default="max-autotune-no-cudagraphs", ) - compile_kwargs["mode"] = mode + compile_kwargs = build_torch_compile_kwargs(mode=mode) logger.info(f"Compiling transformer with mode: {mode}") if self._needs_nvfp4_jit_prewarm(module): @@ -408,8 +402,10 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): prewarm_nvfp4_jit_modules() # TODO(triple-mu): support customized fullgraph and dynamic in the future - module.compile(**compile_kwargs) - self._torch_compiled_module_ids.add(module_id) + self._torch_compile_registry.compile_once( + module, + compile_kwargs=compile_kwargs, + ) def _maybe_enable_cache_dit_and_torch_compile( self, num_inference_steps: int | tuple[int, int], batch: Req diff --git a/python/sglang/multimodal_gen/runtime/server_args.py b/python/sglang/multimodal_gen/runtime/server_args.py index e566b8b81..d8d3dde2c 100644 --- a/python/sglang/multimodal_gen/runtime/server_args.py +++ b/python/sglang/multimodal_gen/runtime/server_args.py @@ -734,6 +734,20 @@ class ServerArgs(DisaggServerArgsMixin): if self.warmup_resolutions is not None: self.warmup = True + if ( + self.enable_torch_compile + and self.warmup_mode is None + and not mode_explicit + and not legacy_explicit + ): + self.warmup = True + self.server_warmup = True + logger.info( + "Automatically enabled server warmup for torch.compile so first " + "real requests do not pay compile latency. Set --warmup-mode off " + "to disable this behavior." + ) + if self.disagg_role != RoleType.MONOLITHIC: self.server_warmup = False @@ -1295,7 +1309,9 @@ class ServerArgs(DisaggServerArgsMixin): "--enable-torch-compile", action=StoreBoolean, default=ServerArgs.enable_torch_compile, - help="Use torch.compile to speed up DiT inference." + help="Use torch.compile to speed up diffusion hot paths. " + + "When no warmup mode is configured, this enables server warmup " + + "so first real requests do not pay compile latency. " + "However, will likely cause precision drifts. See (https://github.com/pytorch/pytorch/issues/145213)", ) parser.add_argument( diff --git a/python/sglang/multimodal_gen/runtime/server_warmup.py b/python/sglang/multimodal_gen/runtime/server_warmup.py index e19e35072..e744665f5 100644 --- a/python/sglang/multimodal_gen/runtime/server_warmup.py +++ b/python/sglang/multimodal_gen/runtime/server_warmup.py @@ -129,9 +129,10 @@ def build_client_warmup_reqs( return_warmup_result=True, server_based_warmup=True, ) - warmup_total = len(warmup_reqs) + warmup_total = sum(1 for req in warmup_reqs if req.is_warmup) for req in warmup_reqs: - req.extra["warmup_total"] = warmup_total + if req.is_warmup: + req.extra["warmup_total"] = warmup_total return warmup_reqs diff --git a/python/sglang/multimodal_gen/runtime/utils/torch_compile.py b/python/sglang/multimodal_gen/runtime/utils/torch_compile.py new file mode 100644 index 000000000..62ca3a178 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/utils/torch_compile.py @@ -0,0 +1,102 @@ +# SPDX-License-Identifier: Apache-2.0 + +import os +from collections.abc import Callable +from dataclasses import dataclass, field +from typing import Any + +import torch.nn as nn + +from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.srt.utils.common import get_compiler_backend + + +def maybe_enable_inductor_compute_comm_overlap() -> None: + try: + import torch._inductor.config as _inductor_cfg + + _inductor_cfg.reorder_for_compute_comm_overlap = True + except ImportError: + pass + + +def build_torch_compile_kwargs(*, mode: str | None) -> dict[str, object]: + compile_kwargs: dict[str, object] = {"fullgraph": False, "dynamic": None} + if current_platform.is_npu(): + compile_kwargs["backend"] = get_compiler_backend() + compile_kwargs["dynamic"] = False + elif mode is not None: + compile_kwargs["mode"] = mode + return compile_kwargs + + +def resolve_torch_compile_mode( + *env_names: str, + config: object | None = None, + default: str, +) -> str: + for env_name in env_names: + mode = os.environ.get(env_name) + if mode: + return mode + mode = getattr(config, "torch_compile_mode", None) + if mode: + return mode + return default + + +@dataclass +class CompiledModuleRegistry: + module_ids: set[int] = field(default_factory=set) + + def is_compiled(self, module: nn.Module) -> bool: + return id(module) in self.module_ids + + def compile_once( + self, + module: nn.Module, + *, + compile_kwargs: dict[str, object], + ) -> bool: + module_id = id(module) + if module_id in self.module_ids: + return False + module.compile(**compile_kwargs) + self.module_ids.add(module_id) + return True + + +class CallableModule(nn.Module): + """Module wrapper for compiling non-forward callables with module.compile""" + + def __init__(self, fn: Callable[..., Any]) -> None: + super().__init__() + self.fn = fn + + def forward(self, *args, **kwargs): + return self.fn(*args, **kwargs) + + +@dataclass +class ActiveTargetCompiledCallable: + """Cache one compiled callable module for the currently active target object""" + + target_id: int | None = None + compiled_module: CallableModule | None = None + + def get_or_compile( + self, + target: object, + fn: Callable[..., Any], + *, + compile_kwargs: dict[str, object], + ) -> Callable[..., Any]: + target_id = id(target) + if self.target_id == target_id and self.compiled_module is not None: + return self.compiled_module + + module = CallableModule(fn) + module.compile(**compile_kwargs) + self.target_id = target_id + self.compiled_module = module + return module diff --git a/python/sglang/multimodal_gen/runtime/warmup_request_builder.py b/python/sglang/multimodal_gen/runtime/warmup_request_builder.py index 14f33022a..4c7245e29 100644 --- a/python/sglang/multimodal_gen/runtime/warmup_request_builder.py +++ b/python/sglang/multimodal_gen/runtime/warmup_request_builder.py @@ -33,6 +33,8 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger logger = init_logger(__name__) DEFAULT_PLACEHOLDER_PROMPT = "warmup" +DEFAULT_DETAILED_PLACEHOLDER_PROMPT = "A detailed image." +TORCH_COMPILE_REAL_PATH_PREWARM_PROMPTS = (DEFAULT_DETAILED_PLACEHOLDER_PROMPT,) DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION = (64, 64) SERVER_WARMUP_IMAGE_FALLBACK_RESOLUTION = (512, 512) SERVER_WARMUP_VIDEO_FALLBACK_RESOLUTION = (832, 480) @@ -260,6 +262,11 @@ def _resolve_warmup_steps( if not server_based_warmup: return warmup_steps + if server_args.enable_torch_compile and server_args.is_arg_explicitly_set( + "warmup_steps" + ): + return warmup_steps + default_steps = sampling_defaults.num_inference_steps if default_steps is None or default_steps <= warmup_steps: return warmup_steps @@ -353,12 +360,29 @@ def build_warmup_reqs( elif negative_prompt is not None and cfg_scale is not None and cfg_scale > 1.0: req_kwargs["do_classifier_free_guidance"] = True - req = Req(**req_kwargs) - 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) + run_real_path_prewarm = server_based_warmup and server_args.enable_torch_compile + prompts = ( + (DEFAULT_PLACEHOLDER_PROMPT,) + TORCH_COMPILE_REAL_PATH_PREWARM_PROMPTS + if run_real_path_prewarm + else (DEFAULT_PLACEHOLDER_PROMPT,) + ) + for prompt_idx, prompt in enumerate(prompts): + prompt_req_kwargs = req_kwargs.copy() + prompt_req_kwargs["prompt"] = prompt + prompt_req_kwargs["sampling_params"] = copy(req_kwargs["sampling_params"]) + req = Req(**prompt_req_kwargs) + if not run_real_path_prewarm or prompt_idx == 0: + req.set_as_warmup(warmup_steps) + else: + req.sampling_params.num_inference_steps = warmup_steps + req.save_output = False + req.suppress_logs = True + req.metrics.suppress_stage_breakdown = True + req.extra["server_internal_prewarm"] = True + 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 diff --git a/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py b/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py index 48e044d2c..d61fb4fa9 100644 --- a/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py +++ b/python/sglang/multimodal_gen/test/unit/test_cfg_parallel_warmup.py @@ -26,6 +26,10 @@ from sglang.multimodal_gen.configs.pipeline_configs.flux_finetuned import ( ) from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import DiffGenerator +from sglang.multimodal_gen.runtime.entrypoints.utils import ( + SetLoraReq, + UnmergeLoraWeightsReq, +) from sglang.multimodal_gen.runtime.managers.scheduler import Scheduler from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( OutputBatch, @@ -37,6 +41,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 format_warmup_req from sglang.multimodal_gen.runtime.warmup_request_builder import ( DEFAULT_PLACEHOLDER_PROMPT, SERVER_WARMUP_IMAGE_FALLBACK_RESOLUTION, @@ -58,7 +63,9 @@ def _make_bare_scheduler(enable_cfg_parallel: bool) -> Scheduler: server_args.warmup_steps = 1 server_args.warmup_resolutions = ["512x512"] server_args.enable_cfg_parallel = enable_cfg_parallel + server_args.enable_torch_compile = False server_args.server_warmup = False + server_args.is_arg_explicitly_set.return_value = False task_type = MagicMock() task_type.requires_image_input.return_value = False @@ -144,6 +151,82 @@ class TestWarmupReqCfgParallel(unittest.TestCase): self.assertIs(req.do_classifier_free_guidance, False) self.assertNotEqual(req.negative_prompt, DEFAULT_PLACEHOLDER_PROMPT) + def test_server_warmup_keeps_minimum_image_steps_without_compile(self): + server_args = _make_bare_scheduler(enable_cfg_parallel=False).server_args + with patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults", + return_value=SamplingParams(num_inference_steps=9), + ): + req = build_warmup_reqs( + server_args, + warmup_resolutions=["512x512"], + server_based_warmup=True, + )[0] + self.assertEqual(req.num_inference_steps, 2) + + def test_torch_compile_respects_explicit_server_warmup_steps(self): + server_args = _make_bare_scheduler(enable_cfg_parallel=False).server_args + server_args.enable_torch_compile = True + server_args.is_arg_explicitly_set.side_effect = lambda name: ( + name == "warmup_steps" + ) + with patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults", + return_value=SamplingParams(num_inference_steps=9), + ): + req = build_warmup_reqs( + server_args, + warmup_resolutions=["512x512"], + server_based_warmup=True, + )[0] + self.assertIn("(512x512, 1/9 steps)", format_warmup_req(req)) + + def test_torch_compile_server_warmup_repeats_each_bucket(self): + server_args = _make_bare_scheduler(enable_cfg_parallel=False).server_args + server_args.enable_torch_compile = True + with patch( + "sglang.multimodal_gen.runtime.warmup_request_builder.get_model_sampling_defaults", + return_value=SamplingParams(num_inference_steps=9), + ): + reqs = build_warmup_reqs( + server_args, + warmup_resolutions=["512x512", "1024x1024"], + server_based_warmup=True, + ) + self.assertEqual(len(reqs), 4) + self.assertEqual( + [(req.width, req.height) for req in reqs], + [(512, 512)] * 2 + [(1024, 1024)] * 2, + ) + self.assertEqual([req.is_warmup for req in reqs], [True, False] * 2) + self.assertEqual([req.num_inference_steps for req in reqs], [2] * 4) + self.assertEqual( + [req.extra.get("server_internal_prewarm", False) for req in reqs], + [False, True] * 2, + ) + self.assertEqual([req.save_output for req in reqs], [False] * 4) + self.assertIsNot(reqs[0].sampling_params, reqs[1].sampling_params) + self.assertEqual(reqs[1].sampling_params.num_inference_steps, 2) + reqs[1].sampling_params.num_inference_steps = 123 + self.assertEqual(reqs[0].sampling_params.num_inference_steps, 2) + + def test_lightweight_warmup_result_ignores_control_requests(self): + scheduler = _make_bare_scheduler(enable_cfg_parallel=False) + + self.assertFalse( + scheduler._should_return_lightweight_warmup_result(SetLoraReq("test")) + ) + self.assertFalse( + scheduler._should_return_lightweight_warmup_result(UnmergeLoraWeightsReq()) + ) + + def test_lightweight_warmup_result_returns_internal_prewarm(self): + scheduler = _make_bare_scheduler(enable_cfg_parallel=False) + req = _make_generation_req() + req.extra["server_internal_prewarm"] = True + + self.assertTrue(scheduler._should_return_lightweight_warmup_result(req)) + def test_req_based_warmup_remains_explicit_legacy_entry(self): scheduler = _make_bare_scheduler(enable_cfg_parallel=False) scheduler.server_args.warmup_resolutions = None @@ -181,6 +264,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args.warmup_resolutions = ["832x480"] server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False task_type = MagicMock() task_type.requires_image_input.return_value = False @@ -217,6 +301,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args = MagicMock() server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False task_type = MagicMock() task_type.requires_image_input.return_value = False @@ -257,6 +342,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args = MagicMock() server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False task_type = MagicMock() task_type.requires_image_input.return_value = False @@ -284,6 +370,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args = MagicMock() server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False task_type = MagicMock() task_type.requires_image_input.return_value = False @@ -325,6 +412,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args = MagicMock() server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False task_type = MagicMock() task_type.requires_image_input.return_value = False @@ -357,6 +445,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args = MagicMock() server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False server_args.backend = "auto" task_type = MagicMock() @@ -385,6 +474,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args = MagicMock() server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False server_args.backend = "diffusers" task_type = MagicMock() @@ -411,6 +501,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args = MagicMock() server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False task_type = MagicMock() task_type.requires_image_input.return_value = False @@ -442,6 +533,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args = MagicMock() server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False task_type = MagicMock() task_type.requires_image_input.return_value = False @@ -481,6 +573,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args = MagicMock() server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False server_args.pipeline_class_name = "LTX2TwoStageHQPipeline" task_type = MagicMock() @@ -515,6 +608,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args = MagicMock() server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False task_type = MagicMock() task_type.requires_image_input.return_value = False @@ -569,6 +663,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args = MagicMock() server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False server_args.pipeline_config.task_type = ModelTaskType.TI2I with patch( @@ -588,6 +683,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args = MagicMock() server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False server_args.pipeline_config.task_type = ModelTaskType.I2I with patch( @@ -607,6 +703,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase): server_args = MagicMock() server_args.warmup_steps = 1 server_args.enable_cfg_parallel = False + server_args.enable_torch_compile = False server_args.pipeline_config.task_type = ModelTaskType.TI2V with patch( diff --git a/python/sglang/multimodal_gen/test/unit/test_decoding_stage_parallelism.py b/python/sglang/multimodal_gen/test/unit/test_decoding_stage_parallelism.py index 8d679d6e8..1ca909efd 100644 --- a/python/sglang/multimodal_gen/test/unit/test_decoding_stage_parallelism.py +++ b/python/sglang/multimodal_gen/test/unit/test_decoding_stage_parallelism.py @@ -2,6 +2,8 @@ import unittest from types import SimpleNamespace from unittest.mock import patch +import torch.nn as nn + from sglang.multimodal_gen.runtime.pipelines_core.stages.base import ( StageParallelismType, ) @@ -11,6 +13,11 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.model_specific_stages.l ) +class FakeVAE(nn.Module): + def decode(self, latents): + return latents + + class TestDecodingStageParallelism(unittest.TestCase): def test_cfg_parallel_uses_replicated_decode_when_decode_group_has_multiple_ranks( self, @@ -37,6 +44,30 @@ class TestDecodingStageParallelism(unittest.TestCase): StageParallelismType.REPLICATED, ) + def test_torch_compile_decode_cache_is_replaced_for_new_vae_instance(self): + vae = FakeVAE() + stage = DecodingStage(vae) + server_args = SimpleNamespace(enable_torch_compile=True) + + with patch( + "torch.compile", + side_effect=lambda fn, **_: fn, + ) as compile_fn: + compiled_vae_decode = stage._get_vae_decode_fn(vae, server_args) + self.assertIs( + stage._get_vae_decode_fn(vae, server_args), compiled_vae_decode + ) + self.assertEqual(compile_fn.call_count, 1) + + new_vae = FakeVAE() + new_compiled_vae_decode = stage._get_vae_decode_fn(new_vae, server_args) + self.assertIsNot(new_compiled_vae_decode, compiled_vae_decode) + self.assertEqual(compile_fn.call_count, 2) + self.assertEqual(stage._compiled_vae_decode.target_id, id(new_vae)) + self.assertIs( + stage._compiled_vae_decode.compiled_module, new_compiled_vae_decode + ) + def test_cfg_parallel_keeps_main_rank_decode_without_parallel_decode(self): stage = object.__new__(DecodingStage) stage.vae = SimpleNamespace(use_parallel_decode=False) diff --git a/python/sglang/multimodal_gen/test/unit/test_server_args.py b/python/sglang/multimodal_gen/test/unit/test_server_args.py index 644686b3c..5f43756dc 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -510,6 +510,7 @@ class TestWarmupModeNormalization(unittest.TestCase): warmup=False, server_warmup=False, warmup_resolutions=None, + enable_torch_compile=False, disagg_role=None, explicit=(), ): @@ -520,6 +521,7 @@ class TestWarmupModeNormalization(unittest.TestCase): sa.warmup = warmup sa.server_warmup = server_warmup sa.warmup_resolutions = warmup_resolutions + sa.enable_torch_compile = enable_torch_compile sa.disagg_role = RoleType.MONOLITHIC if disagg_role is None else disagg_role sa._explicit_arg_names = set(explicit) sa._adjust_warmup() @@ -589,11 +591,39 @@ class TestWarmupModeNormalization(unittest.TestCase): self.assertFalse(sa.server_warmup) self.assertEqual(sa.warmup_mode, "request") + def test_torch_compile_defaults_to_server_warmup(self): + sa = self._resolve(enable_torch_compile=True) + + self.assertEqual(sa.warmup_mode, "server") + self.assertTrue(sa.warmup) + self.assertTrue(sa.server_warmup) + def test_legacy_warmup_on_uses_defaulted_server_mode(self): # `serve --warmup` (legacy ON, mode defaulted to "server" but not # explicit) must resolve to server-based warmup, not silently downgrade # to request mode. sa = self._resolve(warmup_mode="server", warmup=True, explicit=("warmup",)) + + self.assertEqual(sa.warmup_mode, "server") + self.assertTrue(sa.warmup) + self.assertTrue(sa.server_warmup) + + def test_torch_compile_respects_explicit_warmup_off(self): + sa = self._resolve( + warmup_mode="off", + enable_torch_compile=True, + explicit=("warmup_mode",), + ) + self.assertEqual(sa.warmup_mode, "off") + self.assertFalse(sa.warmup) + self.assertFalse(sa.server_warmup) + + def test_torch_compile_uses_server_warmup_for_explicit_resolutions(self): + sa = self._resolve( + warmup_resolutions=["1024x1024"], + enable_torch_compile=True, + explicit=("warmup_resolutions",), + ) self.assertEqual(sa.warmup_mode, "server") self.assertTrue(sa.warmup) self.assertTrue(sa.server_warmup) @@ -624,6 +654,14 @@ class TestWarmupModeNormalization(unittest.TestCase): self.assertFalse(sa.server_warmup) self.assertEqual(sa.warmup_mode, "request") + def test_torch_compile_server_warmup_disabled_for_disagg_role(self): + from sglang.multimodal_gen.runtime.disaggregation.roles import RoleType + + sa = self._resolve(enable_torch_compile=True, disagg_role=RoleType.DENOISER) + self.assertEqual(sa.warmup_mode, "request") + self.assertTrue(sa.warmup) + self.assertFalse(sa.server_warmup) + def test_invalid_mode_raises(self): with self.assertRaises(ValueError): self._resolve(warmup_mode="bogus", explicit=("warmup_mode",)) diff --git a/python/sglang/multimodal_gen/test/unit/test_zimage_pipeline_config.py b/python/sglang/multimodal_gen/test/unit/test_zimage_pipeline_config.py index aac9b99ef..0a5ed7deb 100644 --- a/python/sglang/multimodal_gen/test/unit/test_zimage_pipeline_config.py +++ b/python/sglang/multimodal_gen/test/unit/test_zimage_pipeline_config.py @@ -20,7 +20,9 @@ class TestZImagePipelineConfig(unittest.TestCase): neg_seq_len = 45 batch = SimpleNamespace( prompt_embeds=[torch.ones(pos_seq_len, 2560)], + prompt_seq_lens=[[pos_seq_len]], negative_prompt_embeds=[torch.ones(neg_seq_len, 2560)], + negative_prompt_seq_lens=[[neg_seq_len]], height=16, width=16, )