allow loading AnimateLCM

This commit is contained in:
kijai
2024-03-22 02:40:45 +02:00
parent b3f4405ac7
commit ca6ee5e0e8
2 changed files with 22 additions and 18 deletions
+12 -8
View File
@@ -94,7 +94,7 @@ def ddim_inversion(pipeline, ddim_scheduler, video_latent, num_inv_steps, prompt
def load_weights(
animation_pipeline,
# motion module
motion_module_path = "",
motion_model = "",
motion_module_lora_configs = [],
# domain adapter
adapter_lora_path = "",
@@ -106,20 +106,24 @@ def load_weights(
):
# motion module
unet_state_dict = {}
if isinstance(motion_module_path, str) and motion_module_path != "":
print(f"load motion module from {motion_module_path}")
if motion_module_path.endswith(".safetensors"):
if isinstance(motion_model, str) and motion_model != "":
print(f"load motion module from {motion_model}")
if motion_model.endswith(".safetensors"):
motion_module_state_dict = {}
with safe_open(motion_module_path, framework="pt", device="cpu") as f:
with safe_open(motion_model, 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")
elif motion_model.endswith(".ckpt"):
motion_module_state_dict = torch.load(motion_model, 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", "")
else:
motion_module_state_dict = motion_module_path.model.state_dict()
unet_state_dict = {}
motion_module_state_dict = motion_model.model.state_dict()
if motion_model.model.mm_info.mm_format == "AnimateLCM":
motion_module_state_dict = {k: v for k, v in motion_module_state_dict.items() if "pos_encoder" not in k}
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", "")
+10 -10
View File
@@ -558,11 +558,11 @@ class ADMD_DiffusersLoader:
unet=unet, vae=vae, tokenizer=tokenizer, text_encoder=text_encoder, scheduler=noise_scheduler,
).to(device)
motion_module_path, domain_adapter_path = additional_models
motion_model, domain_adapter_path = additional_models
validation_pipeline = load_weights(
validation_pipeline,
motion_module_path=motion_module_path,
motion_model=motion_model,
adapter_lora_path=domain_adapter_path,
dreambooth_model_path=""
)
@@ -687,11 +687,11 @@ class ADMD_CheckpointLoader:
# Enable gradient checkpointing
unet.enable_gradient_checkpointing()
motion_module_path, domain_adapter_path = additional_models
motion_model, domain_adapter_path = additional_models
validation_pipeline = load_weights(
validation_pipeline,
motion_module_path=motion_module_path,
motion_model=motion_model,
adapter_lora_path=domain_adapter_path,
dreambooth_model_path=""
)
@@ -741,7 +741,7 @@ class ADMD_ComfyModelLoader:
def load_checkpoint(self, model, clip, vae, scheduler, use_xformers, motion_model):
with torch.inference_mode(False):
print(motion_model)
pbar = comfy.utils.ProgressBar(4)
original_config = OmegaConf.load(os.path.join(script_directory, f"configs/v1-inference.yaml"))
ad_unet_config = OmegaConf.load(os.path.join(script_directory, f"configs/ad_unet_config.yaml"))
@@ -823,7 +823,7 @@ class ADMD_ComfyModelLoader:
validation_pipeline = load_weights(
validation_pipeline,
motion_module_path=motion_model,
motion_model=motion_model,
adapter_lora_path=domain_adapter_path,
dreambooth_model_path=""
)
@@ -862,9 +862,9 @@ class ADMD_AdditionalModelSelect:
def select_models(self, motion_module, use_adapter_lora, optional_adapter_lora=""):
additional_models = []
motion_module_path = folder_paths.get_full_path("animatediff_models", motion_module)
if not Path(motion_module_path).is_file():
raise ValueError(f"Motion model {motion_module_path} does not exist")
motion_model = folder_paths.get_full_path("animatediff_models", motion_module)
if not Path(motion_model).is_file():
raise ValueError(f"Motion model {motion_model} does not exist")
if use_adapter_lora:
adapter_lora_path = folder_paths.get_full_path("loras", optional_adapter_lora)
if not Path(adapter_lora_path).is_file():
@@ -872,7 +872,7 @@ class ADMD_AdditionalModelSelect:
else:
adapter_lora_path = ""
additional_models.append(motion_module_path)
additional_models.append(motion_model)
additional_models.append(adapter_lora_path)
return (additional_models,)