From ff08bcdda97a498faaa7583154874be40b733ec8 Mon Sep 17 00:00:00 2001 From: Mick Date: Mon, 7 Sep 2026 09:49:13 +0800 Subject: [PATCH] [diffusion] UX: quiet internal warmup frame searches (#38226) --- .../configs/pipeline_configs/base.py | 2 +- .../configs/pipeline_configs/cosmos3.py | 2 +- .../pipeline_configs/diffusers_generic.py | 2 +- .../configs/pipeline_configs/longlive2.py | 21 ++++--- .../configs/pipeline_configs/mova.py | 15 ++--- .../configs/pipeline_configs/sana_video.py | 2 +- .../configs/pipeline_configs/sana_wm.py | 11 ++-- .../configs/pipeline_configs/wan.py | 9 +-- .../runtime/warmup_request_builder.py | 8 +-- .../unit/test_warmup_probe_degradation.py | 57 +++++++++++++++---- 10 files changed, 86 insertions(+), 43 deletions(-) diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py index 72d70d721..f6180c0a5 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/base.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/base.py @@ -380,7 +380,7 @@ class PipelineConfig: def slice_noise_pred(self, noise, latents): return noise - def adjust_num_frames(self, num_frames): + def adjust_num_frames(self, num_frames, *, log_adjustment: bool = True): return num_frames # tokenize the prompt diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/cosmos3.py b/python/sglang/multimodal_gen/configs/pipeline_configs/cosmos3.py index 38c17b62c..f0321b46d 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/cosmos3.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/cosmos3.py @@ -165,7 +165,7 @@ class Cosmos3Config(PipelineConfig): if self.distilled_sigmas is not None: self.scheduler_class_override = None - def adjust_num_frames(self, num_frames: int) -> int: + def adjust_num_frames(self, num_frames: int, *, log_adjustment: bool = True) -> int: """Round ``num_frames`` so ``(n - 1) % 4 == 0`` for the VAE. Skips rounding when ``num_frames == 1`` (T2I path) so the single diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/diffusers_generic.py b/python/sglang/multimodal_gen/configs/pipeline_configs/diffusers_generic.py index 2a8f146a7..46b560c42 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/diffusers_generic.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/diffusers_generic.py @@ -73,7 +73,7 @@ class DiffusersGenericPipelineConfig(PipelineConfig): """ return width, height - def adjust_num_frames(self, num_frames): + def adjust_num_frames(self, num_frames, *, log_adjustment: bool = True): """ Pass through - diffusers handles frame count. """ diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/longlive2.py b/python/sglang/multimodal_gen/configs/pipeline_configs/longlive2.py index feccd04a2..76318eae4 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/longlive2.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/longlive2.py @@ -37,8 +37,10 @@ class LongLive2T2VConfig(Wan2_2_TI2V_5B_Config): keep_resident_components=("dit", "text_encoder", "vae"), ) - def adjust_num_frames(self, num_frames: int) -> int: - num_frames = super().adjust_num_frames(num_frames) + def adjust_num_frames(self, num_frames: int, *, log_adjustment: bool = True) -> int: + num_frames = super().adjust_num_frames( + num_frames, log_adjustment=log_adjustment + ) vae_scale_factor_temporal = self.vae_config.arch_config.scale_factor_temporal latent_frames = (num_frames - 1) // vae_scale_factor_temporal + 1 block_size = self.dit_config.arch_config.num_frames_per_block @@ -51,13 +53,14 @@ class LongLive2T2VConfig(Wan2_2_TI2V_5B_Config): adjusted_num_frames = ( adjusted_latent_frames - 1 ) * vae_scale_factor_temporal + 1 - logger.warning( - "`num_frames` must map to latent frames divisible by %s for " - "LongLive2 causal denoising. Rounding from %s to %s.", - block_size, - num_frames, - adjusted_num_frames, - ) + if log_adjustment: + logger.warning( + "`num_frames` must map to latent frames divisible by %s for " + "LongLive2 causal denoising. Rounding from %s to %s.", + block_size, + num_frames, + adjusted_num_frames, + ) return adjusted_num_frames def postprocess_image_latent(self, latent_condition, batch): diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/mova.py b/python/sglang/multimodal_gen/configs/pipeline_configs/mova.py index ac5f69009..c83936c3a 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/mova.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/mova.py @@ -112,7 +112,7 @@ class MOVAPipelineConfig(PipelineConfig): ) return image - def adjust_num_frames(self, num_frames: int) -> int: + def adjust_num_frames(self, num_frames: int, *, log_adjustment: bool = True) -> int: if num_frames is None: return num_frames if num_frames % self.time_division_factor != self.time_division_remainder: @@ -122,12 +122,13 @@ class MOVAPipelineConfig(PipelineConfig): * self.time_division_factor + self.time_division_remainder ) - logger.warning( - "`num_frames` (%s) is not compatible with MOVA temporal constraints. " - "Rounding to %s.", - num_frames, - adjusted, - ) + if log_adjustment: + logger.warning( + "`num_frames` (%s) is not compatible with MOVA temporal constraints. " + "Rounding to %s.", + num_frames, + adjusted, + ) return adjusted return num_frames diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/sana_video.py b/python/sglang/multimodal_gen/configs/pipeline_configs/sana_video.py index a1d35266f..0ebc74be7 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/sana_video.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/sana_video.py @@ -63,7 +63,7 @@ class SanaVideoPipelineConfig(PipelineConfig): self.vae_config.load_encoder = False self.vae_config.load_decoder = True - def adjust_num_frames(self, num_frames: int) -> int: + def adjust_num_frames(self, num_frames: int, *, log_adjustment: bool = True) -> int: temporal_scale = self.vae_config.arch_config.temporal_compression_ratio if num_frames < 1: raise ValueError("num_frames must be positive") diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/sana_wm.py b/python/sglang/multimodal_gen/configs/pipeline_configs/sana_wm.py index afd29e82d..1b99d00e9 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/sana_wm.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/sana_wm.py @@ -196,15 +196,16 @@ class SanaWMPipelineConfig(PipelineConfig): return (batch_size, z_dim, T_latent, H_sp, W_sp) - def adjust_num_frames(self, num_frames: int) -> int: + def adjust_num_frames(self, num_frames: int, *, log_adjustment: bool = True) -> int: """Ensure (num_frames - 1) is divisible by VAE temporal stride.""" t_stride = self.vae_stride[0] if (num_frames - 1) % t_stride != 0: adjusted = ((num_frames - 1) // t_stride) * t_stride + 1 - logger.warning( - f"num_frames - 1 must be divisible by temporal stride {t_stride}. " - f"Rounding {num_frames} → {adjusted}." - ) + if log_adjustment: + logger.warning( + f"num_frames - 1 must be divisible by temporal stride {t_stride}. " + f"Rounding {num_frames} → {adjusted}." + ) return adjusted return num_frames diff --git a/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py b/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py index ed37e062c..584fe7269 100644 --- a/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py +++ b/python/sglang/multimodal_gen/configs/pipeline_configs/wan.py @@ -48,12 +48,13 @@ def t5_postprocess_text(outputs: BaseEncoderOutput, _text_inputs) -> torch.Tenso @dataclass class WanI2VCommonConfig(PipelineConfig): # for all wan i2v pipelines - def adjust_num_frames(self, num_frames): + def adjust_num_frames(self, num_frames, *, log_adjustment: bool = True): vae_scale_factor_temporal = self.vae_config.arch_config.scale_factor_temporal if num_frames % vae_scale_factor_temporal != 1: - logger.warning( - f"`num_frames - 1` has to be divisible by {vae_scale_factor_temporal}. Rounding to the nearest number." - ) + if log_adjustment: + logger.warning( + f"`num_frames - 1` has to be divisible by {vae_scale_factor_temporal}. Rounding to the nearest number." + ) num_frames = ( num_frames // vae_scale_factor_temporal * vae_scale_factor_temporal + 1 ) diff --git a/python/sglang/multimodal_gen/runtime/warmup_request_builder.py b/python/sglang/multimodal_gen/runtime/warmup_request_builder.py index 40e49eb53..b6a2a2263 100644 --- a/python/sglang/multimodal_gen/runtime/warmup_request_builder.py +++ b/python/sglang/multimodal_gen/runtime/warmup_request_builder.py @@ -293,11 +293,11 @@ def _lighter_valid_num_frames(server_args: ServerArgs, num_frames: int) -> int: halved count is walked down until it is a fixed point of the contract. """ halved = _halve_num_frames(server_args, num_frames) - adjust = getattr(server_args.pipeline_config, "adjust_num_frames", None) + adjust = server_args.pipeline_config.adjust_num_frames + if halved < adjust(1, log_adjustment=False): + return num_frames for candidate in range(halved, 0, -1): - adjusted = adjust(candidate) if callable(adjust) else candidate - if not isinstance(adjusted, int) or isinstance(adjusted, bool): - adjusted = candidate + adjusted = adjust(candidate, log_adjustment=False) if adjusted == candidate: return candidate if candidate < num_frames else num_frames return num_frames diff --git a/python/sglang/multimodal_gen/test/unit/test_warmup_probe_degradation.py b/python/sglang/multimodal_gen/test/unit/test_warmup_probe_degradation.py index bcbb0ba67..7e2050706 100644 --- a/python/sglang/multimodal_gen/test/unit/test_warmup_probe_degradation.py +++ b/python/sglang/multimodal_gen/test/unit/test_warmup_probe_degradation.py @@ -1,7 +1,9 @@ """A warmup probe that does not fit is retried smaller instead of abandoned.""" from types import SimpleNamespace +from unittest.mock import patch +from sglang.multimodal_gen.configs.pipeline_configs.longlive2 import LongLive2T2VConfig from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.warmup_request_builder import lighten_warmup_req @@ -16,6 +18,7 @@ def _server_args(temporal_compression_ratio: int = 4) -> SimpleNamespace: return SimpleNamespace( pipeline_class_name=None, pipeline_config=SimpleNamespace( + adjust_num_frames=lambda num_frames, **kwargs: num_frames, vae_config=SimpleNamespace(arch_config=arch_config), vae_scale_factor=8, ), @@ -62,17 +65,8 @@ class TestLightenWarmupReq: assert lighten_warmup_req(_server_args(), _req(16, 16, 1)) is None def test_frames_follow_the_model_frame_contract(self): - # LongLive2-style contract: latent frames come in causal blocks of 8, - # so with a temporal ratio of 4 only 29, 61, 93, ... frames are valid. server_args = _server_args() - - def adjust_num_frames(num_frames: int) -> int: - latent = (num_frames - 1) // 4 + 1 - if latent % 8 == 0: - return num_frames - return (max(8, latent // 8 * 8) - 1) * 4 + 1 - - server_args.pipeline_config.adjust_num_frames = adjust_num_frames + server_args.pipeline_config = LongLive2T2VConfig() lighter = lighten_warmup_req(server_args, _req(960, 928, 61)) assert lighter.num_frames == 29 @@ -83,6 +77,30 @@ class TestLightenWarmupReq: assert floor.num_frames == 29 assert floor.width * floor.height <= 960 * 928 // 2 + def test_internal_frame_search_is_quiet_but_user_adjustment_warns(self): + server_args = _server_args() + config = server_args.pipeline_config = LongLive2T2VConfig() + with ( + patch( + "sglang.multimodal_gen.configs.pipeline_configs.wan.logger.warning" + ) as wan_warning, + patch( + "sglang.multimodal_gen.configs.pipeline_configs.longlive2.logger.warning" + ) as causal_warning, + patch.object( + config, "adjust_num_frames", wraps=config.adjust_num_frames + ) as adjust, + ): + lighter = lighten_warmup_req(server_args, _req(1280, 704, 29)) + assert lighter.num_frames == 29 + assert adjust.call_count == 1 + wan_warning.assert_not_called() + causal_warning.assert_not_called() + + assert config.adjust_num_frames(2) == 29 + wan_warning.assert_called_once() + causal_warning.assert_called_once() + def _record(width: int, height: int, num_frames: int, *, peak_gib: float): from sglang.multimodal_gen.runtime.managers.memory_managers.auto_residency import ( @@ -166,6 +184,25 @@ class TestFitAutoResidencyProbe: assert (steps, estimate) == (0, None) assert fitted.num_frames == 81 + def test_longlive_probe_stops_at_the_measured_workload_floor(self): + from sglang.multimodal_gen.runtime.managers.gpu_worker import ( + fit_auto_residency_probe, + ) + + server_args = _server_args() + server_args.pipeline_config = LongLive2T2VConfig() + fitted, _, steps = fit_auto_residency_probe( + _req(1280, 704, 61), + records=[_record(832, 480, 29, peak_gib=30.0)], + free_bytes=8 << 30, + total_bytes=80 << 30, + server_args=server_args, + ) + assert steps > 0 + assert fitted.num_frames == 29 + assert fitted.width > 16 and fitted.height > 16 + assert fitted.width * fitted.height <= 832 * 480 + class TestOutOfMemoryClassification: def test_allocation_failures_from_libraries_count_as_out_of_memory(self):