[diffusion] refactor: validate and document spectrum controls (#33851)

This commit is contained in:
Mick
2026-08-06 23:23:11 +08:00
committed by GitHub
parent 44bde3911a
commit 7195b8e4c7
10 changed files with 143 additions and 9 deletions
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import math
from dataclasses import dataclass
from sglang.multimodal_gen.configs.sample.sampling_params import CacheParams
@@ -67,6 +68,45 @@ class SpectrumParams(CacheParams):
tau_num_steps: int = 50
taylor_order: int = 1
def __post_init__(self) -> None:
finite_numbers = {
"window_size": self.window_size,
"flex_window": self.flex_window,
"w": self.w,
"lam": self.lam,
}
for name, value in finite_numbers.items():
if (
isinstance(value, bool)
or not isinstance(value, (int, float))
or not math.isfinite(value)
):
raise ValueError(f"Spectrum {name} must be a finite number.")
if self.window_size <= 0:
raise ValueError("Spectrum window_size must be greater than zero.")
if self.flex_window < 0:
raise ValueError("Spectrum flex_window must be non-negative.")
if not 0 <= self.w <= 1:
raise ValueError("Spectrum w must be between zero and one.")
if self.lam < 0:
raise ValueError("Spectrum lam must be non-negative.")
non_negative_ints = {"warmup_steps": self.warmup_steps}
positive_ints = {
"m": self.m,
"history_size": self.history_size,
"tau_num_steps": self.tau_num_steps,
}
for name, value in non_negative_ints.items():
if isinstance(value, bool) or not isinstance(value, int) or value < 0:
raise ValueError(f"Spectrum {name} must be a non-negative integer.")
for name, value in positive_ints.items():
if isinstance(value, bool) or not isinstance(value, int) or value <= 0:
raise ValueError(f"Spectrum {name} must be a positive integer.")
if self.taylor_order not in (1, 2, 3):
raise ValueError("Spectrum taylor_order must be one of 1, 2, or 3.")
def get_total_forward_steps(
self, num_inference_steps: int, do_cfg: bool, separate_cfg_branches: bool
) -> int:
@@ -29,6 +29,7 @@ from sglang.multimodal_gen.configs.sample.sampling_params import (
SamplingParams,
_json_safe,
)
from sglang.multimodal_gen.configs.sample.spectrum import SpectrumParams
from sglang.multimodal_gen.configs.sample.teacache import TeaCacheParams
from sglang.multimodal_gen.configs.sample.wan import (
FastWanT2V480PConfig,
@@ -105,6 +106,29 @@ class TestSamplingParamsValidate(unittest.TestCase):
):
SamplingParams(enable_teacache=True, enable_spectrum=True)
def test_spectrum_params_reject_invalid_controls(self):
invalid_controls = (
{"window_size": 0},
{"flex_window": -0.1},
{"w": 1.1},
{"lam": -0.1},
{"warmup_steps": -1},
{"m": 0},
{"history_size": 0},
{"tau_num_steps": 0},
{"taylor_order": 4},
)
for kwargs in invalid_controls:
with self.assertRaises(ValueError):
SpectrumParams(**kwargs)
def test_spectrum_dict_is_validated_when_sampling_params_constructs_it(self):
with self.assertRaisesRegex(ValueError, "history_size"):
SamplingParams(
enable_spectrum=True,
spectrum_params={"history_size": 0},
)
class TestSamplingParamsSubclass(unittest.TestCase):
def test_glm_image_rounds_resolution_up_to_multiple_of_32(self):