WanAnimate block swap fix

This commit is contained in:
kijai
2025-09-19 16:59:57 +03:00
parent 0f1ba64b80
commit e4627466f5
3 changed files with 5 additions and 19 deletions
+1 -13
View File
@@ -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:
+4 -2
View File
@@ -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)
-4
View File
@@ -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