From 29d23e198f690bc2367d9ede95811599cc8eefa3 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E7=8E=8B=E9=B9=A4=E7=94=B7?= Date: Thu, 4 Jun 2026 11:02:23 +0800 Subject: [PATCH] [diffusion] fix: preserve _explicit_fields across dataclasses.replace in DiffGenerator (#25308) --- .../entrypoints/diffusion_generator.py | 6 +++ .../test/unit/test_sampling_params.py | 45 +++++++++++++++++++ 2 files changed, 51 insertions(+) diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py index e9f5de9aa..b80ce7ed4 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py @@ -222,6 +222,12 @@ class DiffGenerator: output_file_name=user_output_file_name, image_path=image_paths_per_prompt[i], ) + # `dataclasses.replace` drops non-field attrs; restore + # `_explicit_fields` so InputValidationStage honors user-supplied + # width/height, and mark the keys overridden above as explicit. + sampling_params._explicit_fields = getattr( + sampling_params_orig, "_explicit_fields", set() + ) | {"prompt", "output_file_name", "image_path"} sampling_params._set_output_file_name() req = prepare_request( server_args=self.server_args, 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 d5b463f8b..386ab3672 100644 --- a/python/sglang/multimodal_gen/test/unit/test_sampling_params.py +++ b/python/sglang/multimodal_gen/test/unit/test_sampling_params.py @@ -279,6 +279,51 @@ class TestSamplingParamsCliArgs(unittest.TestCase): self.assertIn("width", explicit_fields) self.assertIn("height", explicit_fields) + def test_dataclasses_replace_preserves_explicit_fields(self): + """`dataclasses.replace` drops `_explicit_fields`; DiffGenerator must restore it.""" + import dataclasses + + server_args = MagicMock() + server_args.backend = "sglang" + server_args.model_id = None + server_args.pipeline_config = MagicMock() + + with patch.object( + SamplingParams, + "from_pretrained", + side_effect=lambda *args, **kwargs: Flux2SamplingParams(), + ): + sampling_params_orig = SamplingParams.from_user_sampling_params_args( + "dummy-model", + server_args=server_args, + prompt="orig", + image_path="/tmp/in.png", + width=768, + height=512, + ) + + self.assertIn("width", sampling_params_orig._explicit_fields) + self.assertIn("height", sampling_params_orig._explicit_fields) + + cloned = dataclasses.replace( + sampling_params_orig, + prompt="new", + output_file_name=None, + image_path="/tmp/in2.png", + ) + self.assertFalse(hasattr(cloned, "_explicit_fields")) + + # Mirror the restore done in DiffGenerator.generate(). + cloned._explicit_fields = getattr( + sampling_params_orig, "_explicit_fields", set() + ) | {"prompt", "output_file_name", "image_path"} + + explicit = set(cloned.build_request_extra()["explicit_fields"]) + self.assertIn("width", explicit) + self.assertIn("height", explicit) + self.assertIn("prompt", explicit) + self.assertIn("image_path", explicit) + if __name__ == "__main__": unittest.main()