[diffusion] fix: preserve _explicit_fields across dataclasses.replace in DiffGenerator (#25308)

This commit is contained in:
王鹤男
2026-06-04 11:02:23 +08:00
committed by GitHub
parent 11605767e0
commit 29d23e198f
2 changed files with 51 additions and 0 deletions
@@ -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,
@@ -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()