[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)
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,
),
}
@@ -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
@@ -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)
@@ -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
@@ -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
@@ -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(
@@ -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
@@ -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__)
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
@@ -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(
@@ -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)
@@ -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",))
@@ -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,
)