WanAnimate: Encode pose latents in the loop too for better sync and seams between windows

This commit is contained in:
kijai
2025-09-22 02:36:17 +03:00
parent 7ebdcd4abc
commit 02583db554
3 changed files with 48 additions and 36 deletions
+14 -9
View File
@@ -1040,6 +1040,8 @@ class WanVideoAnimateEmbeds:
else:
latent_window_size = latent_window_size + 1
mm.soft_empty_cache()
gc.collect()
vae.to(device)
# Resize and rearrange the input image dimensions
pose_latents = ref_latents = ref_latent = None
@@ -1050,16 +1052,18 @@ class WanVideoAnimateEmbeds:
else:
resized_pose_images = pose_images.permute(3, 0, 1, 2) # C, T, H, W
resized_pose_images = resized_pose_images * 2 - 1
pose_latents = vae.encode([resized_pose_images.to(device, vae.dtype)], device,tiled=tiled_vae)
pose_latents = pose_latents.to(offload_device)
if not looping:
pose_latents = vae.encode([resized_pose_images.to(device, vae.dtype)], device,tiled=tiled_vae)
pose_latents = pose_latents.to(offload_device)
if not looping and pose_latents.shape[2] < latent_window_size:
log.info(f"WanAnimate: Padding pose latents from {pose_latents.shape} to length {latent_window_size}")
pad_len = latent_window_size - pose_latents.shape[2]
pad = torch.zeros(pose_latents.shape[0], pose_latents.shape[1], pad_len, pose_latents.shape[3], pose_latents.shape[4], device=pose_latents.device, dtype=pose_latents.dtype)
pose_latents = torch.cat([pose_latents, pad], dim=2)
print("pose_latents", pose_latents.shape)
del resized_pose_images
if pose_latents.shape[2] < latent_window_size:
log.info(f"WanAnimate: Padding pose latents from {pose_latents.shape} to length {latent_window_size}")
pad_len = latent_window_size - pose_latents.shape[2]
pad = torch.zeros(pose_latents.shape[0], pose_latents.shape[1], pad_len, pose_latents.shape[3], pose_latents.shape[4], device=pose_latents.device, dtype=pose_latents.dtype)
pose_latents = torch.cat([pose_latents, pad], dim=2)
del resized_pose_images
else:
resized_pose_images = resized_pose_images.to(offload_device, dtype=vae.dtype)
bg_latents = None
if bg_images is not None:
@@ -1134,6 +1138,7 @@ class WanVideoAnimateEmbeds:
"negative_clip_context": clip_embeds.get("negative_clip_embeds", None) if clip_embeds is not None else None,
"max_seq_len": seq_len,
"pose_latents": pose_latents,
"pose_images": resized_pose_images if pose_images is not None and looping else None,
"bg_images": resized_bg_images if bg_images is not None and looping else None,
"ref_masks": bg_mask if mask is not None and looping else None,
"ref_latent": ref_latent,
+1 -2
View File
@@ -795,8 +795,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
for r in reader:
all_tensors.extend(r.tensors)
for tensor in all_tensors:
name = tensor.name
name = rename_fuser_block(name)
name = rename_fuser_block(tensor.name)
if "glob" not in name and "audio_proj" in name:
name = name.replace("audio_proj", "multitalk_audio_proj")
load_device = device
+33 -25
View File
@@ -2408,14 +2408,13 @@ class WanVideoSampler:
elif wananimate_loop:
# calculate frame counts
total_frames = num_frames
overlap = 0
refert_num = 1
real_clip_len = frame_window_size - overlap
last_clip_num = (total_frames - overlap) % real_clip_len
real_clip_len = frame_window_size - refert_num
last_clip_num = (total_frames - refert_num) % real_clip_len
extra = 0 if last_clip_num == 0 else real_clip_len - last_clip_num
target_len = total_frames + extra
estimated_iterations = target_len // frame_window_size
estimated_iterations = target_len // real_clip_len
target_latent_len = (target_len - 1) // 4 + estimated_iterations
latent_window_size = (frame_window_size - 1) // 4 + 1
@@ -2425,6 +2424,7 @@ class WanVideoSampler:
ref_images = image_embeds.get("ref_image", None)
ref_masks = image_embeds.get("ref_masks", None)
bg_images = image_embeds.get("bg_images", None)
pose_images = image_embeds.get("pose_images", None)
current_ref_images = face_images = None
@@ -2437,6 +2437,9 @@ class WanVideoSampler:
if bg_images is not None:
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 pose_images is not None:
pose_images_in = tensor_pingpong_pad(pose_images, target_len)
log.info(f"WanAnimate: Pose images {pose_images.shape} padded to shape {pose_images_in.shape}")
# init variables
offloaded = False
@@ -2457,7 +2460,7 @@ class WanVideoSampler:
# outer WanAnimate loop
gen_video_list = []
while True:
if start >= total_frames:
if start + refert_num >= total_frames:
break
mm.soft_empty_cache()
@@ -2478,7 +2481,7 @@ class WanVideoSampler:
if bg_images is not None:
bg_image_slice = bg_images_in[:, start:end].to(device)
else:
bg_image_slice = torch.zeros(3, frame_window_size-mask_reft_len, lat_h * 8, lat_w * 8, device=device, dtype=vae.dtype)
bg_image_slice = torch.zeros(3, frame_window_size, lat_h * 8, lat_w * 8, device=device, dtype=vae.dtype)
if mask_reft_len == 0:
temporal_ref_latents = vae.encode([bg_image_slice], device,tiled=tiled_vae)[0]
else:
@@ -2494,9 +2497,6 @@ class WanVideoSampler:
else:
temporal_ref_latents = temporal_ref_latents[:, :msk.shape[1]]
vae.to(offload_device)
temporal_ref_latents = torch.cat([msk, temporal_ref_latents], dim=0) # 4+C T H W
image_cond_in = torch.cat([ref_latent.to(device), temporal_ref_latents], dim=1) # 4+C T+trefs H W
del temporal_ref_latents, msk, bg_image_slice
@@ -2504,17 +2504,23 @@ class WanVideoSampler:
noise = torch.randn(16, latent_window_size + 1, lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1])
pose_input_slice = None
if wananim_pose_latents is not None:
pose_input_slice = wananim_pose_latents[:, :, start_latent:end_latent].to(device, dtype)
# Pad if slice is too short
if pose_input_slice.shape[2] < latent_window_size:
log.info(f"WanAnimate: Padding pose latents from {pose_input_slice.shape} to length {latent_window_size}")
pad_len = latent_window_size - pose_input_slice.shape[2]
pad = torch.zeros(pose_input_slice.shape[0], pose_input_slice.shape[1], pad_len, pose_input_slice.shape[3], pose_input_slice.shape[4], device=pose_input_slice.device, dtype=pose_input_slice.dtype)
pose_input_slice = torch.cat([pose_input_slice, pad], dim=2)
del pad
pose_input_slice = pose_input_slice.to(device, dtype)
# pose_input_slice = None
# if wananim_pose_latents is not None:
# pose_input_slice = wananim_pose_latents[:, :, start_latent:end_latent].to(device, dtype)
# # Pad if slice is too short
# if pose_input_slice.shape[2] < latent_window_size:
# log.info(f"WanAnimate: Padding pose latents from {pose_input_slice.shape} to length {latent_window_size}")
# pad_len = latent_window_size - pose_input_slice.shape[2]
# pad = torch.zeros(pose_input_slice.shape[0], pose_input_slice.shape[1], pad_len, pose_input_slice.shape[3], pose_input_slice.shape[4], device=pose_input_slice.device, dtype=pose_input_slice.dtype)
# pose_input_slice = torch.cat([pose_input_slice, pad], dim=2)
# del pad
# pose_input_slice = pose_input_slice.to(device, dtype)
if pose_images is not None:
pose_image_slice = pose_images_in[:, start:end].to(device)
print(pose_image_slice.shape)
pose_input_slice = vae.encode([pose_image_slice], device,tiled=tiled_vae).to(dtype)
vae.to(offload_device)
if samples is not None:
input_samples = samples["samples"].squeeze(0).to(noise)
@@ -2631,6 +2637,8 @@ class WanVideoSampler:
videos = vae.decode(latent[:, 1:].unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
del latent
if start != 0:
videos = videos[:, refert_num:]
sampling_pbar.close()
@@ -2646,7 +2654,7 @@ class WanVideoSampler:
videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2)
del cm_result_list
current_ref_images = videos[:, -1:].clone().detach()
current_ref_images = videos[:, -refert_num:].clone().detach()
# optionally save generated samples to disk
if output_path:
@@ -2664,10 +2672,10 @@ class WanVideoSampler:
del videos
iteration_count += 1
start += frame_window_size
end += frame_window_size
start_latent += latent_window_size
end_latent += latent_window_size
start += frame_window_size - refert_num
end += frame_window_size - refert_num
start_latent += latent_window_size - ((refert_num - 1)// 4 + 1)
end_latent += latent_window_size - ((refert_num - 1)// 4 + 1)
if not output_path:
gen_video_samples = torch.cat(gen_video_list, dim=1)