[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,
|
output_file_name=user_output_file_name,
|
||||||
image_path=image_paths_per_prompt[i],
|
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()
|
sampling_params._set_output_file_name()
|
||||||
req = prepare_request(
|
req = prepare_request(
|
||||||
server_args=self.server_args,
|
server_args=self.server_args,
|
||||||
|
|||||||
@@ -279,6 +279,51 @@ class TestSamplingParamsCliArgs(unittest.TestCase):
|
|||||||
self.assertIn("width", explicit_fields)
|
self.assertIn("width", explicit_fields)
|
||||||
self.assertIn("height", 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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user