From b5ac157b41466b5a6654b476ef3ebf23e2fa52d6 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 28 Feb 2024 00:12:34 +0200 Subject: [PATCH] safetensors motion module support --- animatediff/utils/util.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/animatediff/utils/util.py b/animatediff/utils/util.py index 892d334..ece5e8a 100644 --- a/animatediff/utils/util.py +++ b/animatediff/utils/util.py @@ -108,7 +108,13 @@ def load_weights( unet_state_dict = {} if motion_module_path != "": print(f"load motion module from {motion_module_path}") - motion_module_state_dict = torch.load(motion_module_path, map_location="cpu") + if motion_module_path.endswith(".safetensors"): + motion_module_state_dict = {} + with safe_open(motion_module_path, framework="pt", device="cpu") as f: + for key in f.keys(): + motion_module_state_dict[key] = f.get_tensor(key) + elif motion_module_path.endswith(".ckpt"): + motion_module_state_dict = torch.load(motion_module_path, map_location="cpu") motion_module_state_dict = motion_module_state_dict["state_dict"] if "state_dict" in motion_module_state_dict else motion_module_state_dict unet_state_dict.update({name: param for name, param in motion_module_state_dict.items() if "motion_modules." in name}) unet_state_dict.pop("animatediff_config", "")