[diffusion] feat: enable compile warmup for vae decode (#29306)
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user