From 474a851ae3552602951ee6d344309796f69208a9 Mon Sep 17 00:00:00 2001 From: Mick Date: Tue, 17 Mar 2026 08:48:10 +0800 Subject: [PATCH] [diffusion] fix: fix sampling params incorrectly override in cli (#20689) --- .../configs/sample/sampling_params.py | 144 +++++++----------- .../test/unit/test_sampling_params.py | 116 ++++++-------- 2 files changed, 102 insertions(+), 158 deletions(-) diff --git a/python/sglang/multimodal_gen/configs/sample/sampling_params.py b/python/sglang/multimodal_gen/configs/sample/sampling_params.py index 5e9521853..e77fb2f90 100644 --- a/python/sglang/multimodal_gen/configs/sample/sampling_params.py +++ b/python/sglang/multimodal_gen/configs/sample/sampling_params.py @@ -596,213 +596,190 @@ class SamplingParams: @staticmethod def add_cli_args(parser: Any) -> Any: """Add CLI arguments for SamplingParam fields""" - parser.add_argument("--data-type", type=str, nargs="+", default=DataType.VIDEO) - parser.add_argument( + + def add_argument(*name_or_flags, **kwargs): + kwargs.setdefault("default", argparse.SUPPRESS) + return parser.add_argument(*name_or_flags, **kwargs) + + add_argument("--data-type", type=str, nargs="+") + add_argument( "--num-frames-round-down", action="store_true", - default=SamplingParams.num_frames_round_down, ) - parser.add_argument( + add_argument( "--enable-teacache", action="store_true", - default=SamplingParams.enable_teacache, ) # profiling - parser.add_argument( + add_argument( "--profile", action="store_true", - default=SamplingParams.profile, help="Enable torch profiler for denoising stage", ) - parser.add_argument( + add_argument( "--num-profiled-timesteps", type=int, - default=SamplingParams.num_profiled_timesteps, help="Number of timesteps to profile after warmup", ) - parser.add_argument( + add_argument( "--profile-all-stages", action="store_true", dest="profile_all_stages", - default=SamplingParams.profile_all_stages, help="Used with --profile, profile all pipeline stages", ) - parser.add_argument( + add_argument( "--debug", action="store_true", - default=SamplingParams.debug, help="", ) - parser.add_argument( + add_argument( "--prompt", type=str, nargs="+", - default=SamplingParams.prompt, help="Text prompt(s) for generation. Use space-separated values for multiple prompts, e.g., --prompt 'prompt 1' 'prompt 2'", ) - parser.add_argument( + add_argument( "--negative-prompt", type=str, - default=None, help="Negative text prompt for generation", ) - parser.add_argument( + add_argument( "--prompt-path", type=str, - default=SamplingParams.prompt_path, help="Path to a text file containing the prompt", ) - parser.add_argument( + add_argument( "--output-file-name", type=str, - default=SamplingParams.output_file_name, help="Name of the output file", ) - parser.add_argument( + add_argument( "--output-quality", type=str, - default=SamplingParams.output_quality, help="Output quality setting (default, low, medium, high, maximum)", ) - parser.add_argument( + add_argument( "--output-compression", type=int, - default=SamplingParams.output_compression, help="Output compression level (0-100, higher means better quality but larger file size)", ) - parser.add_argument( + add_argument( "--num-outputs-per-prompt", type=int, - default=SamplingParams.num_outputs_per_prompt, help="Number of outputs to generate per prompt", ) - parser.add_argument( + add_argument( "--seed", type=int, - default=SamplingParams.seed, help="Random seed for generation", ) - parser.add_argument( + add_argument( "--generator-device", type=str, - default=SamplingParams.generator_device, choices=["cuda", "musa", "cpu"], help="Device for random generator (cuda, musa or cpu). Default: cuda", ) - parser.add_argument( + add_argument( "--num-frames", type=int, - default=SamplingParams.num_frames, help="Number of frames to generate", ) - parser.add_argument( + add_argument( "--height", type=int, - default=SamplingParams.height, help="Height of generated output", ) - parser.add_argument( + add_argument( "--width", type=int, - default=SamplingParams.width, help="Width of generated output", ) # resolution shortcuts - parser.add_argument( + add_argument( "--4k", action="store_true", dest="resolution_4k", help="Set resolution to 4K (3840x2160)", ) - parser.add_argument( + add_argument( "--2k", action="store_true", dest="resolution_2k", help="Set resolution to 2K (2560x1440)", ) - parser.add_argument( + add_argument( "--1080p", action="store_true", dest="resolution_1080p", help="Set resolution to 1080p (1920x1080)", ) - parser.add_argument( + add_argument( "--720p", action="store_true", dest="resolution_720p", help="Set resolution to 720p (1280x720)", ) - parser.add_argument( + add_argument( "--fps", type=int, - default=SamplingParams.fps, help="Frames per second for saved output", ) - parser.add_argument( + add_argument( "--num-inference-steps", type=int, - default=SamplingParams.num_inference_steps, help="Number of denoising steps", ) - parser.add_argument( + add_argument( "--guidance-scale", type=float, - default=SamplingParams.guidance_scale, help="Classifier-free guidance scale", ) - parser.add_argument( + add_argument( "--guidance-scale-2", type=float, - default=SamplingParams.guidance_scale_2, dest="guidance_scale_2", help="Secondary guidance scale for dual-guidance models (e.g., Wan low-noise expert)", ) - parser.add_argument( + add_argument( "--guidance-rescale", type=float, - default=SamplingParams.guidance_rescale, help="Guidance rescale factor", ) - parser.add_argument( + add_argument( "--cfg-normalization", type=float, - default=SamplingParams.cfg_normalization, # type: ignore[arg-type] dest="cfg_normalization", help=("CFG renormalization factor (for Z-Image). "), ) - parser.add_argument( + add_argument( "--boundary-ratio", type=float, - default=SamplingParams.boundary_ratio, help="Boundary timestep ratio", ) - parser.add_argument( + add_argument( "--save-output", action="store_true", - default=SamplingParams.save_output, help="Whether to save the output to disk", ) - parser.add_argument( + add_argument( "--no-save-output", action="store_false", dest="save_output", help="Don't save the output to disk", ) - parser.add_argument( + add_argument( "--return-frames", action="store_true", - default=SamplingParams.return_frames, help="Whether to return the raw frames", ) - parser.add_argument( + add_argument( "--image-path", type=str, nargs="+", - default=SamplingParams.image_path, help=( "Path(s) to input image(s) for image-to-image / image-to-video " "generation. For multiple images, pass them as space-separated " @@ -810,106 +787,93 @@ class SamplingParams: '--image-path "img1.png" "img2.png"' ), ) - parser.add_argument( + add_argument( "--moba-config-path", type=str, - default=None, help="Path to a JSON file containing V-MoBA specific configurations.", ) - parser.add_argument( + add_argument( "--return-trajectory-latents", action="store_true", - default=SamplingParams.return_trajectory_latents, help="Whether to return the trajectory", ) - parser.add_argument( + add_argument( "--return-trajectory-decoded", action="store_true", - default=SamplingParams.return_trajectory_decoded, help="Whether to return the decoded trajectory", ) - parser.add_argument( + add_argument( "--diffusers-kwargs", type=str, - default=None, help="JSON string of extra kwargs to pass to diffusers pipeline. " 'Example: \'{"output_type": "latent", "clip_skip": 2}\'', ) - parser.add_argument( + add_argument( "--no-override-protected-fields", action="store_true", - default=SamplingParams.no_override_protected_fields, help=( "If set, disallow user params to override fields defined in subclasses." ), ) - parser.add_argument( + add_argument( "--adjust-frames", action=StoreBoolean, - default=SamplingParams.adjust_frames, help=( "Enable/disable adjusting num_frames to evenly split latent frames across GPUs " "and satisfy model temporal constraints. If disabled, tokens might be padded for SP." "Default: true. Examples: --adjust-frames, --adjust-frames true, --adjust-frames false." ), ) - parser.add_argument( + add_argument( "--return-file-paths-only", action=StoreBoolean, - default=SamplingParams.return_file_paths_only, help="If set, output file will be saved early to get a performance boost, while output tensors will not be returned.", ) - parser.add_argument( + add_argument( "--enable-sequence-shard", action=StoreBoolean, - default=SamplingParams.enable_sequence_shard, help="Enable sequence dimension shard with sequence parallelism.", ) - parser.add_argument( + add_argument( "--enable-frame-interpolation", action="store_true", help="Enable post-generation frame interpolation using RIFE 4.22.lite.", ) - parser.add_argument( + add_argument( "--frame-interpolation-exp", type=int, - default=SamplingParams.frame_interpolation_exp, help="Frame interpolation exponent: 1=2x, 2=4x (default: 1).", ) - parser.add_argument( + add_argument( "--frame-interpolation-scale", type=float, - default=SamplingParams.frame_interpolation_scale, help="RIFE inference scale factor (default: 1.0; use 0.5 for high-res).", ) - parser.add_argument( + add_argument( "--frame-interpolation-model-path", type=str, - default=SamplingParams.frame_interpolation_model_path, help="Local directory or HuggingFace repo ID containing RIFE flownet.pkl weights " "(default: elfgum/RIFE-4.22.lite, downloaded automatically). " "Only RIFE 4.22.lite architecture is supported; other RIFE versions or " "frame interpolation models are not compatible.", ) - parser.add_argument( + add_argument( "--enable-upscaling", action="store_true", help="Enable post-generation upscaling using Real-ESRGAN.", ) - parser.add_argument( + add_argument( "--upscaling-model-path", type=str, - default=SamplingParams.upscaling_model_path, help="Local .pth file, HuggingFace repo ID, or repo_id:filename for Real-ESRGAN weights " "(default: ai-forever/Real-ESRGAN with RealESRGAN_x4.pth). " "Only RRDBNet (e.g. RealESRGAN_x4plus) and SRVGGNetCompact (e.g. realesr-animevideov3) " "architectures are supported; other super-resolution models are not compatible. " "Use 'repo_id:filename' to specify a custom weight file from a HF repo.", ) - parser.add_argument( + add_argument( "--upscaling-scale", type=int, - default=SamplingParams.upscaling_scale, help="Upscaling factor (default: 4).", ) return parser diff --git a/python/sglang/multimodal_gen/test/unit/test_sampling_params.py b/python/sglang/multimodal_gen/test/unit/test_sampling_params.py index 90a805696..722d4d12a 100644 --- a/python/sglang/multimodal_gen/test/unit/test_sampling_params.py +++ b/python/sglang/multimodal_gen/test/unit/test_sampling_params.py @@ -6,6 +6,7 @@ from sglang.multimodal_gen.configs.sample.diffusers_generic import ( DiffusersGenericSamplingParams, ) from sglang.multimodal_gen.configs.sample.flux import FluxSamplingParams +from sglang.multimodal_gen.configs.sample.qwenimage import QwenImageSamplingParams from sglang.multimodal_gen.configs.sample.sampling_params import SamplingParams @@ -68,81 +69,60 @@ class TestSamplingParamsSubclass(unittest.TestCase): DiffusersGenericSamplingParams(num_frames=0) -class TestNegativePromptMerge(unittest.TestCase): - """Regression tests for negative_prompt not being passed through CLI""" +class TestSamplingParamsCliArgs(unittest.TestCase): + def _parse_cli_kwargs(self, argv: list[str]) -> dict: + parser = argparse.ArgumentParser() + SamplingParams.add_cli_args(parser) + args = parser.parse_args(argv) + return SamplingParams.get_cli_args(args) - def test_get_cli_args_filters_none(self): - ns = argparse.Namespace(negative_prompt=None, prompt="hello") - result = SamplingParams.get_cli_args(ns) - self.assertNotIn("negative_prompt", result) - self.assertEqual(result["prompt"], "hello") + def _make_qwen_image_params(self, argv: list[str]) -> QwenImageSamplingParams: + return QwenImageSamplingParams(**self._parse_cli_kwargs(argv)) - def test_get_cli_args_keeps_explicit_value(self): - ns = argparse.Namespace(negative_prompt="ugly, blurry") - result = SamplingParams.get_cli_args(ns) - self.assertEqual(result["negative_prompt"], "ugly, blurry") + def test_get_cli_args_drops_unset_sampling_params(self): + self.assertEqual(self._parse_cli_kwargs([]), {}) - def test_merge_preserves_subclass_default_when_not_explicit(self): - """Without explicit_fields, value matching base default is not merged, - so the subclass default (empty string) is preserved.""" + def test_get_cli_args_keeps_explicit_sampling_params(self): + kwargs = self._parse_cli_kwargs( + [ + "--guidance-scale", + str(SamplingParams.guidance_scale), + "--negative-prompt", + SamplingParams.negative_prompt, + "--save-output", + ] + ) + + self.assertEqual(kwargs["guidance_scale"], SamplingParams.guidance_scale) + self.assertEqual(kwargs["negative_prompt"], SamplingParams.negative_prompt) + self.assertTrue(kwargs["save_output"]) + + def test_qwen_image_cli_path_preserves_model_defaults(self): + params = self._make_qwen_image_params([]) + + self.assertEqual(params.negative_prompt, " ") + self.assertEqual(params.guidance_scale, 4.0) + + def test_qwen_image_cli_path_allows_explicit_override_to_base_defaults(self): + params = self._make_qwen_image_params( + [ + "--guidance-scale", + str(SamplingParams.guidance_scale), + "--negative-prompt", + SamplingParams.negative_prompt, + ] + ) + + self.assertEqual(params.guidance_scale, SamplingParams.guidance_scale) + self.assertEqual(params.negative_prompt, SamplingParams.negative_prompt) + + def test_merge_allows_explicit_field_matching_base_default(self): target = DiffusersGenericSamplingParams() - self.assertEqual(target.negative_prompt, "") + user = SamplingParams(negative_prompt=SamplingParams.negative_prompt) - user = SamplingParams() - target._merge_with_user_params(user) - self.assertEqual(target.negative_prompt, "") - - def test_merge_applies_different_negative_prompt(self): - target = DiffusersGenericSamplingParams() - user = SamplingParams(negative_prompt="ugly, blurry") - target._merge_with_user_params(user) - self.assertEqual(target.negative_prompt, "ugly, blurry") - - def test_merge_explicit_field_matching_base_default(self): - """Even when the user value matches the base-class default, it should - still be applied if listed in explicit_fields.""" - base_default = SamplingParams.negative_prompt - target = DiffusersGenericSamplingParams() - self.assertEqual(target.negative_prompt, "") - - user = SamplingParams(negative_prompt=base_default) target._merge_with_user_params(user, explicit_fields={"negative_prompt"}) - self.assertEqual(target.negative_prompt, base_default) - def test_cli_roundtrip_no_negative_prompt(self): - """Simulate CLI without --negative-prompt: subclass default is kept.""" - ns = argparse.Namespace(negative_prompt=None, width=512, height=512) - kwargs = SamplingParams.get_cli_args(ns) - self.assertNotIn("negative_prompt", kwargs) - - user = SamplingParams(**kwargs) - target = DiffusersGenericSamplingParams() - target._merge_with_user_params(user, explicit_fields=set(kwargs.keys())) - self.assertEqual(target.negative_prompt, "") - - def test_cli_roundtrip_with_negative_prompt(self): - """Simulate CLI with --negative-prompt: user value is applied.""" - user_neg = "bad quality, watermark" - ns = argparse.Namespace(negative_prompt=user_neg, width=512, height=512) - kwargs = SamplingParams.get_cli_args(ns) - user = SamplingParams(**kwargs) - - target = DiffusersGenericSamplingParams() - target._merge_with_user_params(user, explicit_fields=set(kwargs.keys())) - self.assertEqual(target.negative_prompt, user_neg) - - def test_cli_roundtrip_with_base_default_negative_prompt(self): - """Simulate CLI where --negative-prompt value matches the base default: - user value should still be applied (not dropped).""" - base_default = SamplingParams.negative_prompt - ns = argparse.Namespace(negative_prompt=base_default, width=512, height=512) - kwargs = SamplingParams.get_cli_args(ns) - self.assertIn("negative_prompt", kwargs) - - user = SamplingParams(**kwargs) - target = DiffusersGenericSamplingParams() - target._merge_with_user_params(user, explicit_fields=set(kwargs.keys())) - self.assertEqual(target.negative_prompt, base_default) + self.assertEqual(target.negative_prompt, SamplingParams.negative_prompt) if __name__ == "__main__":