[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)
|
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,
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user