From daa638befb6c6cb49b5547994c977494174db474 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 19 Aug 2025 01:04:26 +0300 Subject: [PATCH] fix UniAnimate --- nodes.py | 4 ++-- nodes_model_loading.py | 6 +++--- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/nodes.py b/nodes.py index 8699795..5ccb396 100644 --- a/nodes.py +++ b/nodes.py @@ -1954,8 +1954,8 @@ class WanVideoSampler: # UniAnimate if unianimate_poses is not None: - transformer.dwpose_embedding.to(device, model["dtype"]) - dwpose_data = unianimate_poses["pose"].to(device, model["dtype"]) + transformer.dwpose_embedding.to(device, dtype) + dwpose_data = unianimate_poses["pose"].to(device, dtype) dwpose_data = torch.cat([dwpose_data[:,:,:1].repeat(1,1,3,1,1), dwpose_data], dim=2) dwpose_data = transformer.dwpose_embedding(dwpose_data) log.info(f"UniAnimate pose embed shape: {dwpose_data.shape}") diff --git a/nodes_model_loading.py b/nodes_model_loading.py index ac1133a..1b39dd0 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -748,7 +748,7 @@ def load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_dev block_idx = None #print("block_idx:", block_idx) #print("vace_block_idx:", vace_block_idx) - if "loras" in name: + if "loras" in name or "dwpose" in name or "randomref" in name: continue dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else weight_dtype dtype_to_use = weight_dtype if sd[name.replace("_orig_mod.", "")].dtype == weight_dtype else dtype_to_use @@ -863,7 +863,7 @@ def add_lora_weights(patcher, lora, base_dtype, merge_loras=False): if isinstance(lora_strength, list): if merge_loras: raise ValueError("LoRA strength should be a single value when merge_loras=True") - transformer.lora_scheduling_enabled = True + patcher.model.diffusion_model.lora_scheduling_enabled = True if lora_strength == 0: log.warning(f"LoRA {lora_path} has strength 0, skipping...") continue @@ -871,7 +871,7 @@ def add_lora_weights(patcher, lora, base_dtype, merge_loras=False): if "dwpose_embedding.0.weight" in lora_sd: #unianimate from .unianimate.nodes import update_transformer log.info("Unianimate LoRA detected, patching model...") - transformer = update_transformer(transformer, lora_sd) + patcher.model.diffusion_model = update_transformer(patcher.model.diffusion_model, lora_sd) lora_sd = standardize_lora_key_format(lora_sd)