[diffusion] UX: quiet internal warmup frame searches (#38226)
This commit is contained in:
@@ -380,7 +380,7 @@ class PipelineConfig:
|
|||||||
def slice_noise_pred(self, noise, latents):
|
def slice_noise_pred(self, noise, latents):
|
||||||
return noise
|
return noise
|
||||||
|
|
||||||
def adjust_num_frames(self, num_frames):
|
def adjust_num_frames(self, num_frames, *, log_adjustment: bool = True):
|
||||||
return num_frames
|
return num_frames
|
||||||
|
|
||||||
# tokenize the prompt
|
# tokenize the prompt
|
||||||
|
|||||||
@@ -165,7 +165,7 @@ class Cosmos3Config(PipelineConfig):
|
|||||||
if self.distilled_sigmas is not None:
|
if self.distilled_sigmas is not None:
|
||||||
self.scheduler_class_override = 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.
|
"""Round ``num_frames`` so ``(n - 1) % 4 == 0`` for the VAE.
|
||||||
|
|
||||||
Skips rounding when ``num_frames == 1`` (T2I path) so the single
|
Skips rounding when ``num_frames == 1`` (T2I path) so the single
|
||||||
|
|||||||
@@ -73,7 +73,7 @@ class DiffusersGenericPipelineConfig(PipelineConfig):
|
|||||||
"""
|
"""
|
||||||
return width, height
|
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.
|
Pass through - diffusers handles frame count.
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -37,8 +37,10 @@ class LongLive2T2VConfig(Wan2_2_TI2V_5B_Config):
|
|||||||
keep_resident_components=("dit", "text_encoder", "vae"),
|
keep_resident_components=("dit", "text_encoder", "vae"),
|
||||||
)
|
)
|
||||||
|
|
||||||
def adjust_num_frames(self, num_frames: int) -> int:
|
def adjust_num_frames(self, num_frames: int, *, log_adjustment: bool = True) -> int:
|
||||||
num_frames = super().adjust_num_frames(num_frames)
|
num_frames = super().adjust_num_frames(
|
||||||
|
num_frames, log_adjustment=log_adjustment
|
||||||
|
)
|
||||||
vae_scale_factor_temporal = self.vae_config.arch_config.scale_factor_temporal
|
vae_scale_factor_temporal = self.vae_config.arch_config.scale_factor_temporal
|
||||||
latent_frames = (num_frames - 1) // vae_scale_factor_temporal + 1
|
latent_frames = (num_frames - 1) // vae_scale_factor_temporal + 1
|
||||||
block_size = self.dit_config.arch_config.num_frames_per_block
|
block_size = self.dit_config.arch_config.num_frames_per_block
|
||||||
@@ -51,6 +53,7 @@ class LongLive2T2VConfig(Wan2_2_TI2V_5B_Config):
|
|||||||
adjusted_num_frames = (
|
adjusted_num_frames = (
|
||||||
adjusted_latent_frames - 1
|
adjusted_latent_frames - 1
|
||||||
) * vae_scale_factor_temporal + 1
|
) * vae_scale_factor_temporal + 1
|
||||||
|
if log_adjustment:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"`num_frames` must map to latent frames divisible by %s for "
|
"`num_frames` must map to latent frames divisible by %s for "
|
||||||
"LongLive2 causal denoising. Rounding from %s to %s.",
|
"LongLive2 causal denoising. Rounding from %s to %s.",
|
||||||
|
|||||||
@@ -112,7 +112,7 @@ class MOVAPipelineConfig(PipelineConfig):
|
|||||||
)
|
)
|
||||||
return image
|
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:
|
if num_frames is None:
|
||||||
return num_frames
|
return num_frames
|
||||||
if num_frames % self.time_division_factor != self.time_division_remainder:
|
if num_frames % self.time_division_factor != self.time_division_remainder:
|
||||||
@@ -122,6 +122,7 @@ class MOVAPipelineConfig(PipelineConfig):
|
|||||||
* self.time_division_factor
|
* self.time_division_factor
|
||||||
+ self.time_division_remainder
|
+ self.time_division_remainder
|
||||||
)
|
)
|
||||||
|
if log_adjustment:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"`num_frames` (%s) is not compatible with MOVA temporal constraints. "
|
"`num_frames` (%s) is not compatible with MOVA temporal constraints. "
|
||||||
"Rounding to %s.",
|
"Rounding to %s.",
|
||||||
|
|||||||
@@ -63,7 +63,7 @@ class SanaVideoPipelineConfig(PipelineConfig):
|
|||||||
self.vae_config.load_encoder = False
|
self.vae_config.load_encoder = False
|
||||||
self.vae_config.load_decoder = True
|
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
|
temporal_scale = self.vae_config.arch_config.temporal_compression_ratio
|
||||||
if num_frames < 1:
|
if num_frames < 1:
|
||||||
raise ValueError("num_frames must be positive")
|
raise ValueError("num_frames must be positive")
|
||||||
|
|||||||
@@ -196,11 +196,12 @@ class SanaWMPipelineConfig(PipelineConfig):
|
|||||||
|
|
||||||
return (batch_size, z_dim, T_latent, H_sp, W_sp)
|
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."""
|
"""Ensure (num_frames - 1) is divisible by VAE temporal stride."""
|
||||||
t_stride = self.vae_stride[0]
|
t_stride = self.vae_stride[0]
|
||||||
if (num_frames - 1) % t_stride != 0:
|
if (num_frames - 1) % t_stride != 0:
|
||||||
adjusted = ((num_frames - 1) // t_stride) * t_stride + 1
|
adjusted = ((num_frames - 1) // t_stride) * t_stride + 1
|
||||||
|
if log_adjustment:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"num_frames - 1 must be divisible by temporal stride {t_stride}. "
|
f"num_frames - 1 must be divisible by temporal stride {t_stride}. "
|
||||||
f"Rounding {num_frames} → {adjusted}."
|
f"Rounding {num_frames} → {adjusted}."
|
||||||
|
|||||||
@@ -48,9 +48,10 @@ def t5_postprocess_text(outputs: BaseEncoderOutput, _text_inputs) -> torch.Tenso
|
|||||||
@dataclass
|
@dataclass
|
||||||
class WanI2VCommonConfig(PipelineConfig):
|
class WanI2VCommonConfig(PipelineConfig):
|
||||||
# for all wan i2v pipelines
|
# 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
|
vae_scale_factor_temporal = self.vae_config.arch_config.scale_factor_temporal
|
||||||
if num_frames % vae_scale_factor_temporal != 1:
|
if num_frames % vae_scale_factor_temporal != 1:
|
||||||
|
if log_adjustment:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
f"`num_frames - 1` has to be divisible by {vae_scale_factor_temporal}. Rounding to the nearest number."
|
f"`num_frames - 1` has to be divisible by {vae_scale_factor_temporal}. Rounding to the nearest number."
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -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 count is walked down until it is a fixed point of the contract.
|
||||||
"""
|
"""
|
||||||
halved = _halve_num_frames(server_args, num_frames)
|
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):
|
for candidate in range(halved, 0, -1):
|
||||||
adjusted = adjust(candidate) if callable(adjust) else candidate
|
adjusted = adjust(candidate, log_adjustment=False)
|
||||||
if not isinstance(adjusted, int) or isinstance(adjusted, bool):
|
|
||||||
adjusted = candidate
|
|
||||||
if adjusted == candidate:
|
if adjusted == candidate:
|
||||||
return candidate if candidate < num_frames else num_frames
|
return candidate if candidate < num_frames else num_frames
|
||||||
return num_frames
|
return num_frames
|
||||||
|
|||||||
@@ -1,7 +1,9 @@
|
|||||||
"""A warmup probe that does not fit is retried smaller instead of abandoned."""
|
"""A warmup probe that does not fit is retried smaller instead of abandoned."""
|
||||||
|
|
||||||
from types import SimpleNamespace
|
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.configs.sample.sampling_params import SamplingParams
|
||||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||||
from sglang.multimodal_gen.runtime.warmup_request_builder import lighten_warmup_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(
|
return SimpleNamespace(
|
||||||
pipeline_class_name=None,
|
pipeline_class_name=None,
|
||||||
pipeline_config=SimpleNamespace(
|
pipeline_config=SimpleNamespace(
|
||||||
|
adjust_num_frames=lambda num_frames, **kwargs: num_frames,
|
||||||
vae_config=SimpleNamespace(arch_config=arch_config),
|
vae_config=SimpleNamespace(arch_config=arch_config),
|
||||||
vae_scale_factor=8,
|
vae_scale_factor=8,
|
||||||
),
|
),
|
||||||
@@ -62,17 +65,8 @@ class TestLightenWarmupReq:
|
|||||||
assert lighten_warmup_req(_server_args(), _req(16, 16, 1)) is None
|
assert lighten_warmup_req(_server_args(), _req(16, 16, 1)) is None
|
||||||
|
|
||||||
def test_frames_follow_the_model_frame_contract(self):
|
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()
|
server_args = _server_args()
|
||||||
|
server_args.pipeline_config = LongLive2T2VConfig()
|
||||||
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
|
|
||||||
|
|
||||||
lighter = lighten_warmup_req(server_args, _req(960, 928, 61))
|
lighter = lighten_warmup_req(server_args, _req(960, 928, 61))
|
||||||
assert lighter.num_frames == 29
|
assert lighter.num_frames == 29
|
||||||
@@ -83,6 +77,30 @@ class TestLightenWarmupReq:
|
|||||||
assert floor.num_frames == 29
|
assert floor.num_frames == 29
|
||||||
assert floor.width * floor.height <= 960 * 928 // 2
|
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):
|
def _record(width: int, height: int, num_frames: int, *, peak_gib: float):
|
||||||
from sglang.multimodal_gen.runtime.managers.memory_managers.auto_residency import (
|
from sglang.multimodal_gen.runtime.managers.memory_managers.auto_residency import (
|
||||||
@@ -166,6 +184,25 @@ class TestFitAutoResidencyProbe:
|
|||||||
assert (steps, estimate) == (0, None)
|
assert (steps, estimate) == (0, None)
|
||||||
assert fitted.num_frames == 81
|
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:
|
class TestOutOfMemoryClassification:
|
||||||
def test_allocation_failures_from_libraries_count_as_out_of_memory(self):
|
def test_allocation_failures_from_libraries_count_as_out_of_memory(self):
|
||||||
|
|||||||
Reference in New Issue
Block a user