From 64cbd28e000466a8b5952aa1c0d93fba6d6595a2 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E5=B0=8F=E5=85=AD=E5=A6=9E=E5=84=BF?= Date: Sat, 27 Dec 2025 16:04:00 +0800 Subject: [PATCH] feat(wanvideo): add manual start reference support for WanAnimate loop --- nodes.py | 12 ++++++++++-- nodes_sampler.py | 13 ++++++++++--- 2 files changed, 20 insertions(+), 5 deletions(-) diff --git a/nodes.py b/nodes.py index 67f5142..b91cd39 100644 --- a/nodes.py +++ b/nodes.py @@ -1135,6 +1135,7 @@ class WanVideoAnimateEmbeds: "face_images": ("IMAGE", {"tooltip": "end frame"}), "bg_images": ("IMAGE", {"tooltip": "background images"}), "mask": ("MASK", {"tooltip": "mask"}), + "start_ref_image": ("IMAGE", {"tooltip": "start ref image"}), "tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}), } } @@ -1145,7 +1146,7 @@ class WanVideoAnimateEmbeds: CATEGORY = "WanVideoWrapper" def process(self, vae, width, height, num_frames, force_offload, frame_window_size, colormatch, pose_strength, face_strength, - ref_images=None, pose_images=None, face_images=None, clip_embeds=None, tiled_vae=False, bg_images=None, mask=None): + ref_images=None, pose_images=None, face_images=None, clip_embeds=None, tiled_vae=False, bg_images=None, mask=None, start_ref_image=None): W = (width // 16) * 16 H = (height // 16) * 16 @@ -1156,7 +1157,7 @@ class WanVideoAnimateEmbeds: num_refs = ref_images.shape[0] if ref_images is not None else 0 num_frames = ((num_frames - 1) // 4) * 4 + 1 - looping = num_frames > frame_window_size + looping = num_frames > frame_window_size or start_ref_image is not None if num_frames < frame_window_size: frame_window_size = num_frames @@ -1254,6 +1255,12 @@ class WanVideoAnimateEmbeds: resized_face_images = (resized_face_images * 2 - 1).unsqueeze(0) resized_face_images = resized_face_images.to(offload_device, dtype=vae.dtype) + if start_ref_image is not None: + if start_ref_image.shape[1] != H or start_ref_image.shape[2] != W: + resized_start_ref_image = common_upscale(start_ref_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1) + else: + resized_start_ref_image = start_ref_image.permute(3, 0, 1, 2) # C, T, H, W + resized_start_ref_image = resized_start_ref_image[:3] * 2 - 1 seq_len = math.ceil((target_shape[2] * target_shape[3]) / 4 * target_shape[1]) @@ -1273,6 +1280,7 @@ class WanVideoAnimateEmbeds: "is_masked": mask is not None, "ref_latent": ref_latent, "ref_image": resized_ref_images if ref_images is not None else None, + "start_ref_image": resized_start_ref_image if start_ref_image is not None else None, "face_pixels": resized_face_images if face_images is not None else None, "num_frames": num_frames, "target_shape": target_shape, diff --git a/nodes_sampler.py b/nodes_sampler.py index 57b73e8..fc1f6c3 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -2178,7 +2178,11 @@ class WanVideoSampler: bg_images = image_embeds.get("bg_images", None) pose_images = image_embeds.get("pose_images", None) - current_ref_images = face_images = face_images_in = None + current_ref_images = image_embeds.get("start_ref_image", None) + if current_ref_images is not None: + log.info( + "WanAnimate: Detected manual start reference image, enabling continuous generation across windows.") + face_images = face_images_in = None if wananim_face_pixels is not None: face_images = tensor_pingpong_pad(wananim_face_pixels, target_len) @@ -2217,7 +2221,10 @@ class WanVideoSampler: mm.soft_empty_cache() - mask_reft_len = 0 if start == 0 else refert_num + if current_ref_images is not None: + mask_reft_len = refert_num + else: + mask_reft_len = 0 if start == 0 else refert_num self.cache_state = [None, None] @@ -2396,7 +2403,7 @@ 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: + if start != 0 or current_ref_images is not None: videos = videos[:, refert_num:] sampling_pbar.close()