[diffusion] feat: enable compile warmup for vae decode (#29306)

This commit is contained in:
Mick
2026-07-04 15:25:26 +08:00
committed by GitHub
parent 576dc31e33
commit 03962d4238
13 changed files with 428 additions and 43 deletions
@@ -202,6 +202,12 @@ class ZImagePipelineConfig(ZImageRolloutPipelineMixin, ImagePipelineConfig):
plan = self._build_zimage_sp_plan(batch) plan = self._build_zimage_sp_plan(batch)
return plan 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): def _split_text_embeds_for_dit(self, batch, *, negative: bool = False):
"""Return per-request text tensors, trimming padded batched embeddings.""" """Return per-request text tensors, trimming padded batched embeddings."""
embeds = batch.negative_prompt_embeds if negative else batch.prompt_embeds embeds = batch.negative_prompt_embeds if negative else batch.prompt_embeds
@@ -217,7 +223,7 @@ class ZImagePipelineConfig(ZImageRolloutPipelineMixin, ImagePipelineConfig):
return embeds return embeds
if embeds.ndim == 2: if embeds.ndim == 2:
return [embeds] return [self._pad_text_embed_for_dit(embeds)]
if embeds.ndim != 3: if embeds.ndim != 3:
raise ValueError( raise ValueError(
@@ -231,7 +237,8 @@ class ZImagePipelineConfig(ZImageRolloutPipelineMixin, ImagePipelineConfig):
expected_batch_size=int(embeds.shape[0]), expected_batch_size=int(embeds.shape[0]),
) )
return [ 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): 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 if get_sp_world_size() > 1
else None 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): 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 if get_sp_world_size() > 1
else None else None
), ),
"caption_valid_lens": torch.tensor(
self.require_text_seq_lens(batch, 0, negative=use_negative_embeds),
device=device,
dtype=torch.long,
),
} }
@@ -57,6 +57,7 @@ from sglang.multimodal_gen.runtime.server_args import (
) )
from sglang.multimodal_gen.runtime.server_warmup import ( from sglang.multimodal_gen.runtime.server_warmup import (
SchedulerWarmupMixin, SchedulerWarmupMixin,
get_first_generation_req,
is_warmup_req, is_warmup_req,
should_return_warmup_result, should_return_warmup_result,
) )
@@ -571,6 +572,12 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
) -> List[OutputBatch]: ) -> List[OutputBatch]:
return [OutputBatch(error=error_msg) for _ in reqs] 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( def return_result(
self, self,
output_batch: OutputBatch, output_batch: OutputBatch,
@@ -1048,8 +1055,11 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
is_warmup = is_warmup_req(processed_req) is_warmup = is_warmup_req(processed_req)
self._log_warmup_result(output_batch, processed_req, is_warmup) self._log_warmup_result(output_batch, processed_req, is_warmup)
if is_warmup and should_return_warmup_result(processed_req): should_return_lightweight_warmup_result = (
# only keep the necessary lightweight payloads 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() output_batch.drop_payload_for_warmup()
self.return_result( self.return_result(
output_batch, identity, should_not_return=False output_batch, identity, should_not_return=False
@@ -817,6 +817,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
patch_size: int, patch_size: int,
f_patch_size: int, f_patch_size: int,
image_seq_len_target: int | None = None, image_seq_len_target: int | None = None,
caption_valid_lens: torch.Tensor | None = None,
): ):
"""Patchify images and pad image/caption tokens to batch targets. """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
) )
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_ori_len = cap_feat.size(0)
cap_padding_len = cap_seq_len_target - cap_ori_len cap_padding_len = cap_seq_len_target - cap_ori_len
cap_padded_feat = torch.cat( cap_padded_feat = torch.cat(
@@ -854,7 +860,10 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
dim=0, dim=0,
) )
all_cap_feats_out.append(cap_padded_feat) 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 target_image_seq_len = image_seq_len_target or 0
for image in all_image: for image in all_image:
@@ -885,12 +894,15 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
all_image_size.append(image_size) all_image_size.append(image_size)
all_image_valid_lens.append(image_ori_len) 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 ( return (
torch.stack(all_image_out, dim=0), torch.stack(all_image_out, dim=0),
torch.stack(all_cap_feats_out, dim=0), torch.stack(all_cap_feats_out, dim=0),
all_image_size, all_image_size,
all_image_valid_lens, all_image_valid_lens,
all_cap_valid_lens, cap_valid_lens_out,
) )
@staticmethod @staticmethod
@@ -923,12 +935,16 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
@staticmethod @staticmethod
def _replace_padding_with_token( def _replace_padding_with_token(
tensor: torch.Tensor, tensor: torch.Tensor,
valid_lens: list[int], valid_lens: list[int] | torch.Tensor,
pad_token: torch.Tensor, pad_token: torch.Tensor,
) -> torch.Tensor: ) -> torch.Tensor:
"""Replace padded token rows after each valid sequence length.""" """Replace padded token rows after each valid sequence length."""
positions = torch.arange(tensor.shape[1], device=tensor.device).unsqueeze(0) 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 pad_mask = positions >= lengths
if pad_mask.any(): if pad_mask.any():
tensor = tensor.clone() tensor = tensor.clone()
@@ -945,6 +961,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
f_patch_size=1, f_patch_size=1,
freqs_cis=None, freqs_cis=None,
image_seq_len_target: int | None = None, image_seq_len_target: int | None = None,
caption_valid_lens: torch.Tensor | None = None,
**kwargs, **kwargs,
): ):
assert patch_size in self.all_patch_size assert patch_size in self.all_patch_size
@@ -968,6 +985,7 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
patch_size, patch_size,
f_patch_size, f_patch_size,
image_seq_len_target=image_seq_len_target, 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) x, _ = self.all_x_embedder[f"{patch_size}-{f_patch_size}"](x)
@@ -8,6 +8,7 @@ Decoding stage for diffusion pipelines.
import weakref import weakref
import torch import torch
import torch.nn as nn
from sglang.multimodal_gen.runtime.distributed import ( from sglang.multimodal_gen.runtime.distributed import (
get_decode_parallel_world_size, get_decode_parallel_world_size,
@@ -35,6 +36,11 @@ from sglang.multimodal_gen.runtime.utils.precision import (
resolve_precision, resolve_precision,
temporary_module_dtype, 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__) logger = init_logger(__name__)
@@ -105,6 +111,7 @@ class DecodingStage(PipelineStage):
self.vae: ParallelTiledVAE = vae self.vae: ParallelTiledVAE = vae
self.pipeline = weakref.ref(pipeline) if pipeline else None self.pipeline = weakref.ref(pipeline) if pipeline else None
self.component_name = component_name self.component_name = component_name
self._compiled_vae_decode = ActiveTargetCompiledCallable()
def component_uses( def component_uses(
self, server_args: ServerArgs, stage_name: str | None = None 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): def scale_and_shift(self, latents: torch.Tensor, server_args):
return scale_and_shift_latents(latents, server_args, self.vae) 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() @torch.no_grad()
def decode( def decode(
self, self,
@@ -209,7 +242,7 @@ class DecodingStage(PipelineStage):
with temporary_module_dtype( with temporary_module_dtype(
self.vae, vae_dtype, enabled=should_cast_vae self.vae, vae_dtype, enabled=should_cast_vae
) as 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) image = _ensure_tensor_decode_output(decode_output)
# De-normalize image to [0, 1] range # De-normalize image to [0, 1] range
@@ -7,7 +7,6 @@ Denoising stage for diffusion pipelines.
import inspect import inspect
import math import math
import os
import time import time
import weakref import weakref
from collections.abc import Callable from collections.abc import Callable
@@ -114,8 +113,13 @@ from sglang.multimodal_gen.runtime.utils.precision import (
resolve_precision, resolve_precision,
) )
from sglang.multimodal_gen.runtime.utils.profiler import SGLDiffusionProfiler 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.multimodal_gen.utils import dict_to_3d_list
from sglang.srt.utils.common import get_compiler_backend
logger = init_logger(__name__) logger = init_logger(__name__)
@@ -209,7 +213,7 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
# cache-dit state (for delayed mounting and idempotent control) # cache-dit state (for delayed mounting and idempotent control)
self._cache_dit_enabled = False self._cache_dit_enabled = False
self._cached_num_steps = None 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 hidden_size = self.server_args.pipeline_config.dit_config.hidden_size
num_attention_heads = ( num_attention_heads = (
@@ -373,31 +377,21 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
if envs.SGLANG_CACHE_DIT_ENABLED and not self._cache_dit_enabled: if envs.SGLANG_CACHE_DIT_ENABLED and not self._cache_dit_enabled:
logger.debug("Deferring torch.compile until cache-dit is enabled") logger.debug("Deferring torch.compile until cache-dit is enabled")
return return
module_id = id(module) if self._torch_compile_registry.is_compiled(module):
if module_id in self._torch_compiled_module_ids:
return return
compile_kwargs: dict[str, Any] = {"fullgraph": False, "dynamic": None}
if current_platform.is_npu(): if current_platform.is_npu():
backend = get_compiler_backend() compile_kwargs = build_torch_compile_kwargs(mode=None)
compile_kwargs["backend"] = backend
compile_kwargs["dynamic"] = False
logger.info("Compiling transformer with torchair backend on NPU") logger.info("Compiling transformer with torchair backend on NPU")
else: else:
try: maybe_enable_inductor_compute_comm_overlap()
import torch._inductor.config as _inductor_cfg
_inductor_cfg.reorder_for_compute_comm_overlap = True
except ImportError:
pass
dit_config = getattr(self.server_args.pipeline_config, "dit_config", None) dit_config = getattr(self.server_args.pipeline_config, "dit_config", None)
mode = os.environ.get("SGLANG_TORCH_COMPILE_MODE") or getattr( mode = resolve_torch_compile_mode(
dit_config, "SGLANG_TORCH_COMPILE_MODE",
"torch_compile_mode", config=dit_config,
"max-autotune-no-cudagraphs", 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}") logger.info(f"Compiling transformer with mode: {mode}")
if self._needs_nvfp4_jit_prewarm(module): if self._needs_nvfp4_jit_prewarm(module):
@@ -408,8 +402,10 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
prewarm_nvfp4_jit_modules() prewarm_nvfp4_jit_modules()
# TODO(triple-mu): support customized fullgraph and dynamic in the future # TODO(triple-mu): support customized fullgraph and dynamic in the future
module.compile(**compile_kwargs) self._torch_compile_registry.compile_once(
self._torch_compiled_module_ids.add(module_id) module,
compile_kwargs=compile_kwargs,
)
def _maybe_enable_cache_dit_and_torch_compile( def _maybe_enable_cache_dit_and_torch_compile(
self, num_inference_steps: int | tuple[int, int], batch: Req self, num_inference_steps: int | tuple[int, int], batch: Req
@@ -734,6 +734,20 @@ class ServerArgs(DisaggServerArgsMixin):
if self.warmup_resolutions is not None: if self.warmup_resolutions is not None:
self.warmup = True 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: if self.disagg_role != RoleType.MONOLITHIC:
self.server_warmup = False self.server_warmup = False
@@ -1295,7 +1309,9 @@ class ServerArgs(DisaggServerArgsMixin):
"--enable-torch-compile", "--enable-torch-compile",
action=StoreBoolean, action=StoreBoolean,
default=ServerArgs.enable_torch_compile, 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)", + "However, will likely cause precision drifts. See (https://github.com/pytorch/pytorch/issues/145213)",
) )
parser.add_argument( parser.add_argument(
@@ -129,9 +129,10 @@ def build_client_warmup_reqs(
return_warmup_result=True, return_warmup_result=True,
server_based_warmup=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: for req in warmup_reqs:
req.extra["warmup_total"] = warmup_total if req.is_warmup:
req.extra["warmup_total"] = warmup_total
return warmup_reqs return warmup_reqs
@@ -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
@@ -33,6 +33,8 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
logger = init_logger(__name__) logger = init_logger(__name__)
DEFAULT_PLACEHOLDER_PROMPT = "warmup" 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) DEFAULT_LIGHTWEIGHT_IMAGE_RESOLUTION = (64, 64)
SERVER_WARMUP_IMAGE_FALLBACK_RESOLUTION = (512, 512) SERVER_WARMUP_IMAGE_FALLBACK_RESOLUTION = (512, 512)
SERVER_WARMUP_VIDEO_FALLBACK_RESOLUTION = (832, 480) SERVER_WARMUP_VIDEO_FALLBACK_RESOLUTION = (832, 480)
@@ -260,6 +262,11 @@ def _resolve_warmup_steps(
if not server_based_warmup: if not server_based_warmup:
return warmup_steps 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 default_steps = sampling_defaults.num_inference_steps
if default_steps is None or default_steps <= warmup_steps: if default_steps is None or default_steps <= warmup_steps:
return 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: 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_kwargs["do_classifier_free_guidance"] = True
req = Req(**req_kwargs) run_real_path_prewarm = server_based_warmup and server_args.enable_torch_compile
req.set_as_warmup(warmup_steps) prompts = (
if return_warmup_result: (DEFAULT_PLACEHOLDER_PROMPT,) + TORCH_COMPILE_REAL_PATH_PREWARM_PROMPTS
req.extra["return_warmup_result"] = True if run_real_path_prewarm
if server_based_warmup: else (DEFAULT_PLACEHOLDER_PROMPT,)
req.extra["server_based_warmup"] = True )
warmup_reqs.append(req) 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 return warmup_reqs
@@ -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.configs.sample.sampling_params import SamplingParams
from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import DiffGenerator 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.managers.scheduler import Scheduler
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import (
OutputBatch, 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 ( from sglang.multimodal_gen.runtime.pipelines_core.stages.input_validation import (
InputValidationStage, InputValidationStage,
) )
from sglang.multimodal_gen.runtime.server_warmup import format_warmup_req
from sglang.multimodal_gen.runtime.warmup_request_builder import ( from sglang.multimodal_gen.runtime.warmup_request_builder import (
DEFAULT_PLACEHOLDER_PROMPT, DEFAULT_PLACEHOLDER_PROMPT,
SERVER_WARMUP_IMAGE_FALLBACK_RESOLUTION, 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_steps = 1
server_args.warmup_resolutions = ["512x512"] server_args.warmup_resolutions = ["512x512"]
server_args.enable_cfg_parallel = enable_cfg_parallel server_args.enable_cfg_parallel = enable_cfg_parallel
server_args.enable_torch_compile = False
server_args.server_warmup = False server_args.server_warmup = False
server_args.is_arg_explicitly_set.return_value = False
task_type = MagicMock() task_type = MagicMock()
task_type.requires_image_input.return_value = False 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.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_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): def test_req_based_warmup_remains_explicit_legacy_entry(self):
scheduler = _make_bare_scheduler(enable_cfg_parallel=False) scheduler = _make_bare_scheduler(enable_cfg_parallel=False)
scheduler.server_args.warmup_resolutions = None scheduler.server_args.warmup_resolutions = None
@@ -181,6 +264,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args.warmup_resolutions = ["832x480"] server_args.warmup_resolutions = ["832x480"]
server_args.warmup_steps = 1 server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock() task_type = MagicMock()
task_type.requires_image_input.return_value = False task_type.requires_image_input.return_value = False
@@ -217,6 +301,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args = MagicMock() server_args = MagicMock()
server_args.warmup_steps = 1 server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock() task_type = MagicMock()
task_type.requires_image_input.return_value = False task_type.requires_image_input.return_value = False
@@ -257,6 +342,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args = MagicMock() server_args = MagicMock()
server_args.warmup_steps = 1 server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock() task_type = MagicMock()
task_type.requires_image_input.return_value = False task_type.requires_image_input.return_value = False
@@ -284,6 +370,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args = MagicMock() server_args = MagicMock()
server_args.warmup_steps = 1 server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock() task_type = MagicMock()
task_type.requires_image_input.return_value = False task_type.requires_image_input.return_value = False
@@ -325,6 +412,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args = MagicMock() server_args = MagicMock()
server_args.warmup_steps = 1 server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock() task_type = MagicMock()
task_type.requires_image_input.return_value = False task_type.requires_image_input.return_value = False
@@ -357,6 +445,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args = MagicMock() server_args = MagicMock()
server_args.warmup_steps = 1 server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
server_args.backend = "auto" server_args.backend = "auto"
task_type = MagicMock() task_type = MagicMock()
@@ -385,6 +474,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args = MagicMock() server_args = MagicMock()
server_args.warmup_steps = 1 server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
server_args.backend = "diffusers" server_args.backend = "diffusers"
task_type = MagicMock() task_type = MagicMock()
@@ -411,6 +501,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args = MagicMock() server_args = MagicMock()
server_args.warmup_steps = 1 server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock() task_type = MagicMock()
task_type.requires_image_input.return_value = False task_type.requires_image_input.return_value = False
@@ -442,6 +533,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args = MagicMock() server_args = MagicMock()
server_args.warmup_steps = 1 server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock() task_type = MagicMock()
task_type.requires_image_input.return_value = False task_type.requires_image_input.return_value = False
@@ -481,6 +573,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args = MagicMock() server_args = MagicMock()
server_args.warmup_steps = 1 server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
server_args.pipeline_class_name = "LTX2TwoStageHQPipeline" server_args.pipeline_class_name = "LTX2TwoStageHQPipeline"
task_type = MagicMock() task_type = MagicMock()
@@ -515,6 +608,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args = MagicMock() server_args = MagicMock()
server_args.warmup_steps = 1 server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
task_type = MagicMock() task_type = MagicMock()
task_type.requires_image_input.return_value = False task_type.requires_image_input.return_value = False
@@ -569,6 +663,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args = MagicMock() server_args = MagicMock()
server_args.warmup_steps = 1 server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
server_args.pipeline_config.task_type = ModelTaskType.TI2I server_args.pipeline_config.task_type = ModelTaskType.TI2I
with patch( with patch(
@@ -588,6 +683,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args = MagicMock() server_args = MagicMock()
server_args.warmup_steps = 1 server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
server_args.pipeline_config.task_type = ModelTaskType.I2I server_args.pipeline_config.task_type = ModelTaskType.I2I
with patch( with patch(
@@ -607,6 +703,7 @@ class TestWarmupReqCfgParallel(unittest.TestCase):
server_args = MagicMock() server_args = MagicMock()
server_args.warmup_steps = 1 server_args.warmup_steps = 1
server_args.enable_cfg_parallel = False server_args.enable_cfg_parallel = False
server_args.enable_torch_compile = False
server_args.pipeline_config.task_type = ModelTaskType.TI2V server_args.pipeline_config.task_type = ModelTaskType.TI2V
with patch( with patch(
@@ -2,6 +2,8 @@ import unittest
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import patch from unittest.mock import patch
import torch.nn as nn
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import ( from sglang.multimodal_gen.runtime.pipelines_core.stages.base import (
StageParallelismType, 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): class TestDecodingStageParallelism(unittest.TestCase):
def test_cfg_parallel_uses_replicated_decode_when_decode_group_has_multiple_ranks( def test_cfg_parallel_uses_replicated_decode_when_decode_group_has_multiple_ranks(
self, self,
@@ -37,6 +44,30 @@ class TestDecodingStageParallelism(unittest.TestCase):
StageParallelismType.REPLICATED, 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): def test_cfg_parallel_keeps_main_rank_decode_without_parallel_decode(self):
stage = object.__new__(DecodingStage) stage = object.__new__(DecodingStage)
stage.vae = SimpleNamespace(use_parallel_decode=False) stage.vae = SimpleNamespace(use_parallel_decode=False)
@@ -510,6 +510,7 @@ class TestWarmupModeNormalization(unittest.TestCase):
warmup=False, warmup=False,
server_warmup=False, server_warmup=False,
warmup_resolutions=None, warmup_resolutions=None,
enable_torch_compile=False,
disagg_role=None, disagg_role=None,
explicit=(), explicit=(),
): ):
@@ -520,6 +521,7 @@ class TestWarmupModeNormalization(unittest.TestCase):
sa.warmup = warmup sa.warmup = warmup
sa.server_warmup = server_warmup sa.server_warmup = server_warmup
sa.warmup_resolutions = warmup_resolutions 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.disagg_role = RoleType.MONOLITHIC if disagg_role is None else disagg_role
sa._explicit_arg_names = set(explicit) sa._explicit_arg_names = set(explicit)
sa._adjust_warmup() sa._adjust_warmup()
@@ -589,11 +591,39 @@ class TestWarmupModeNormalization(unittest.TestCase):
self.assertFalse(sa.server_warmup) self.assertFalse(sa.server_warmup)
self.assertEqual(sa.warmup_mode, "request") 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): def test_legacy_warmup_on_uses_defaulted_server_mode(self):
# `serve --warmup` (legacy ON, mode defaulted to "server" but not # `serve --warmup` (legacy ON, mode defaulted to "server" but not
# explicit) must resolve to server-based warmup, not silently downgrade # explicit) must resolve to server-based warmup, not silently downgrade
# to request mode. # to request mode.
sa = self._resolve(warmup_mode="server", warmup=True, explicit=("warmup",)) 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.assertEqual(sa.warmup_mode, "server")
self.assertTrue(sa.warmup) self.assertTrue(sa.warmup)
self.assertTrue(sa.server_warmup) self.assertTrue(sa.server_warmup)
@@ -624,6 +654,14 @@ class TestWarmupModeNormalization(unittest.TestCase):
self.assertFalse(sa.server_warmup) self.assertFalse(sa.server_warmup)
self.assertEqual(sa.warmup_mode, "request") 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): def test_invalid_mode_raises(self):
with self.assertRaises(ValueError): with self.assertRaises(ValueError):
self._resolve(warmup_mode="bogus", explicit=("warmup_mode",)) self._resolve(warmup_mode="bogus", explicit=("warmup_mode",))
@@ -20,7 +20,9 @@ class TestZImagePipelineConfig(unittest.TestCase):
neg_seq_len = 45 neg_seq_len = 45
batch = SimpleNamespace( batch = SimpleNamespace(
prompt_embeds=[torch.ones(pos_seq_len, 2560)], 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_embeds=[torch.ones(neg_seq_len, 2560)],
negative_prompt_seq_lens=[[neg_seq_len]],
height=16, height=16,
width=16, width=16,
) )