From e4627466f501d48344620b4d4d385b8abedaea48 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 19 Sep 2025 16:59:57 +0300 Subject: [PATCH] WanAnimate block swap fix --- nodes.py | 14 +------------- nodes_model_loading.py | 6 ++++-- wanvideo/modules/model.py | 4 ---- 3 files changed, 5 insertions(+), 19 deletions(-) diff --git a/nodes.py b/nodes.py index fe39659..d80b6b3 100644 --- a/nodes.py +++ b/nodes.py @@ -4294,9 +4294,7 @@ class WanVideoSampler: bg_images = image_embeds.get("bg_images", None) pose_input_latents = current_ref_images = face_images = None - #if wananim_pose_latents is not None: - #pose_input_latents = tensor_pingpong_pad(wananim_pose_latents, target_latent_len) - #log.info(f"WanAnimate: Pose input {wananim_pose_latents.shape} padded to shape {pose_input_latents.shape}") + if wananim_face_pixels is not None: face_images = tensor_pingpong_pad(wananim_face_pixels, target_len) log.info(f"WanAnimate: Face input {wananim_face_pixels.shape} padded to shape {face_images.shape}") @@ -4307,11 +4305,6 @@ class WanVideoSampler: bg_images_in = tensor_pingpong_pad(bg_images, target_len) log.info(f"WanAnimate: BG images {bg_images.shape} padded to shape {bg_images.shape}") - # if replace_flag: - # bg_images, mask_images = self.prepare_source_for_replace(src_bg_path, src_mask_path) - # bg_images = inputs_padding(bg_images, target_len) - # mask_images = inputs_padding(mask_images, target_len) - # init variables offloaded = False @@ -4344,11 +4337,6 @@ class WanVideoSampler: vae.to(device) if ref_masks is not None: msk = ref_masks_in[:, start_latent:end_latent].to(device, dtype) - # if msk.shape[1] < latent_window_size: - # log.info(f"WanAnimate: Padding ref masks from {msk.shape} to length {latent_window_size}") - # pad_length = latent_window_size - msk.shape[1] - # last_frame = msk[:, -1:].repeat(1, pad_length, 1, 1) - # msk = torch.cat([msk, last_frame], dim=1) else: msk = torch.zeros(4, latent_window_size, lat_h, lat_w, device=device, dtype=dtype) if bg_images is not None: diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 81d1827..ac57738 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -792,7 +792,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None, vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0]) except Exception: vace_block_idx = None - elif "blocks." in name: + elif "blocks." in name and "face" not in name: try: block_idx = int(name.split("blocks.")[1].split(".")[0]) except Exception: @@ -831,7 +831,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None, vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0]) except Exception: vace_block_idx = None - elif "blocks." in name: + elif "blocks." in name and "face" not in name: try: block_idx = int(name.split("blocks.")[1].split(".")[0]) except Exception: @@ -868,6 +868,8 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None, cnt += 1 if cnt % 100 == 0: pbar.update(100) + #for name, param in transformer.named_parameters(): + # print(name, param.device, param.dtype) pbar.update_absolute(0) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index e8ff190..f678ab9 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -2242,7 +2242,6 @@ class WanModel(torch.nn.Module): ip_img_ids[:, :, :, 2] = ip_img_ids[:, :, :, 2] + torch.linspace(w_len + freq_offset, w_len + freq_offset + w_ip - 1, steps=w_ip, device=x.device, dtype=x.dtype).reshape(1, 1, -1) ip_img_ids = repeat(ip_img_ids, "t h w c -> b (t h w) c", b=1) freqs_ip = self.rope_embedder(ip_img_ids).movedim(1, 2) - #print("freqs_ip shape:", freqs_ip.shape) # EchoShot cross attn freqs inner_c = None @@ -2256,9 +2255,6 @@ class WanModel(torch.nn.Module): x = torch.cat([x, motion_encoded], dim=1) freqs = torch.cat([freqs, freqs_motion], dim=1) - #t = torch.repeat_interleave(t, 2, dim=1) - #t = torch.cat([t, torch.zeros((t.shape[0], 3), device=t.device, dtype=t.dtype)], dim=1) - # time embeddings if t.dim() == 2: b, f = t.shape