[diffusion] chore: set allowing overriding protected fields of sampling params as default behavior (#14471)

This commit is contained in:
Mick
2025-12-06 00:22:42 +08:00
committed by GitHub
parent 662809874c
commit a89045603b
5 changed files with 49 additions and 62 deletions
@@ -138,7 +138,7 @@ class SamplingParams:
return_trajectory_latents: bool = False # returns all latents for each timestep
return_trajectory_decoded: bool = False # returns decoded latents for each timestep
# if True, allow user params to override subclass-defined protected fields
override_protected_fields: bool = False
no_override_protected_fields: bool = True
# whether to adjust num_frames for multi-GPU friendly splitting (default: True)
adjust_frames: bool = True
@@ -290,15 +290,6 @@ class SamplingParams:
self._set_output_file_name()
self.log(server_args=server_args)
def update(self, source_dict: dict[str, Any]) -> None:
for key, value in source_dict.items():
if hasattr(self, key):
setattr(self, key, value)
else:
logger.exception("%s has no attribute %s", type(self).__name__, key)
self.__post_init__()
@classmethod
def from_pretrained(cls, model_path: str, **kwargs) -> "SamplingParams":
from sglang.multimodal_gen.registry import get_model_info
@@ -522,12 +513,11 @@ class SamplingParams:
help="Whether to return the decoded trajectory",
)
parser.add_argument(
"--override-protected-fields",
"--no-override-protected-fields",
action="store_true",
default=SamplingParams.override_protected_fields,
default=SamplingParams.no_override_protected_fields,
help=(
"If set, allow user params to override fields defined in subclasses "
"(protected by default)."
"If set, disallow user params to override fields defined in subclasses."
),
)
parser.add_argument(
@@ -583,32 +573,21 @@ class SamplingParams:
subclass_defined_fields = set(type(self).__annotations__.keys())
# global switch: if True, allow overriding protected fields
allow_override_protected = bool(
user_params.override_protected_fields or self.override_protected_fields
)
# Compare against current instance to avoid constructing a default instance
default_params = SamplingParams()
allow_override_protected = not user_params.no_override_protected_fields
for field in dataclasses.fields(user_params):
field_name = field.name
user_value = getattr(user_params, field_name)
default_value = getattr(default_params, field_name)
default_value = getattr(self, field_name)
# A field is considered user-modified if its value is different from
# the default, with an exception for `output_file_name` which is
# auto-generated with a random component.
is_user_modified = (
user_value != default_value
if field_name != "output_file_name"
else user_params.output_file_path is not None
)
# the default
is_user_modified = user_value != default_value
if is_user_modified and (
allow_override_protected or field_name not in subclass_defined_fields
):
if hasattr(self, field_name):
setattr(self, field_name, user_value)
self.height_not_provided = user_params.height_not_provided
self.width_not_provided = user_params.width_not_provided
self.__post_init__()