[diffusion] fix: preserve _explicit_fields across dataclasses.replace in DiffGenerator (#25308)
This commit is contained in:
@@ -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()
|
||||
|
||||
Reference in New Issue
Block a user