[diffusion] CI: enable UT (#20690)

This commit is contained in:
Mick
2026-03-17 07:44:04 +08:00
committed by GitHub
parent 2ccdb7373e
commit 1eea744855
5 changed files with 95 additions and 21 deletions
+16 -12
View File
@@ -35,21 +35,25 @@ logger = init_logger(__name__)
T = TypeVar("T")
def _expand_path_value(field_name: str, value: Any) -> Any:
eu = os.path.expanduser
if field_name.endswith("_path") and isinstance(value, str):
return eu(value)
if field_name.endswith("_path") and isinstance(value, list):
return [eu(x) if isinstance(x, str) else x for x in value]
if field_name.endswith("_paths") and isinstance(value, dict):
return {k: eu(p) if isinstance(p, str) else p for k, p in value.items()}
return value
def expand_path_kwargs(kwargs: dict[str, Any]) -> dict[str, Any]:
return {key: _expand_path_value(key, value) for key, value in kwargs.items()}
def expand_path_fields(obj) -> None:
"""In-place expanduser on all dataclass fields whose name ends with '_path' or '_paths'."""
eu = os.path.expanduser
for f in fields(obj):
v = getattr(obj, f.name)
if f.name.endswith("_path") and isinstance(v, str):
setattr(obj, f.name, eu(v))
elif f.name.endswith("_path") and isinstance(v, list):
setattr(obj, f.name, [eu(x) if isinstance(x, str) else x for x in v])
elif f.name.endswith("_paths") and isinstance(v, dict):
setattr(
obj,
f.name,
{k: eu(p) if isinstance(p, str) else p for k, p in v.items()},
)
setattr(obj, f.name, _expand_path_value(f.name, getattr(obj, f.name)))
# TODO(will): used to convert server_args.precision to torch.dtype. Find a