[diffusion] chore: validate sampling params (#16677)
Co-authored-by: root <root@huchong2-0.huchong2.podvm.svc.cluster.local>
This commit is contained in:
@@ -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": [
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user