[diffusion] chore: validate sampling params (#16677)

Co-authored-by: root <root@huchong2-0.huchong2.podvm.svc.cluster.local>
This commit is contained in:
Hu Chong
2026-01-12 16:00:10 +08:00
committed by GitHub
parent c54c70ab62
commit f44c63eef7
3 changed files with 128 additions and 3 deletions

View File

@@ -27,6 +27,8 @@ SUITES = {
"test_lora_format_adapter.py",
# cli test
"../cli/test_generate_t2i_perf.py",
# unit tests (no server needed)
"../test_sampling_params_validate.py",
# add new 1-gpu test files here
],
"2-gpu": [

View File

@@ -0,0 +1,49 @@
import math
import unittest
from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams
class TestSamplingParamsValidate(unittest.TestCase):
def test_prompt_path_suffix(self):
with self.assertRaisesRegex(ValueError, r"prompt_path"):
SamplingParams(prompt_path="bad.png")
def test_num_outputs_per_prompt_must_be_positive(self):
with self.assertRaisesRegex(ValueError, r"num_outputs_per_prompt"):
SamplingParams(num_outputs_per_prompt=0)
def test_fps_must_be_positive_int(self):
with self.assertRaisesRegex(ValueError, r"\bfps\b"):
SamplingParams(fps=0)
with self.assertRaisesRegex(ValueError, r"\bfps\b"):
SamplingParams(fps=None) # type: ignore[arg-type]
def test_num_inference_steps_optional_but_if_set_must_be_positive(self):
SamplingParams(num_inference_steps=None)
with self.assertRaisesRegex(ValueError, r"num_inference_steps"):
SamplingParams(num_inference_steps=-1)
def test_guidance_scale_must_be_finite_non_negative_if_set(self):
SamplingParams(guidance_scale=None)
with self.assertRaisesRegex(ValueError, r"guidance_scale"):
SamplingParams(guidance_scale=math.nan)
with self.assertRaisesRegex(ValueError, r"guidance_scale"):
SamplingParams(guidance_scale=-0.1)
def test_guidance_rescale_must_be_finite_non_negative(self):
with self.assertRaisesRegex(ValueError, r"guidance_rescale"):
SamplingParams(guidance_rescale=-1.0)
with self.assertRaisesRegex(ValueError, r"guidance_rescale"):
SamplingParams(guidance_rescale=math.inf)
def test_boundary_ratio_range(self):
SamplingParams(boundary_ratio=None)
with self.assertRaisesRegex(ValueError, r"boundary_ratio"):
SamplingParams(boundary_ratio=1.5)
with self.assertRaisesRegex(ValueError, r"boundary_ratio"):
SamplingParams(boundary_ratio=math.nan)
if __name__ == "__main__":
unittest.main()