From f85ea5a48a52d009901029c8abc521bd8cd0b119 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 5 Dec 2025 13:05:35 +0200 Subject: [PATCH] add "empty_frame_pad_image" input to I2V encode node for easier SVI lora use This pads the empty frames with the given image, as should be done with SVI-shot and SVI 2.0 LoRAs --- nodes.py | 77 +++++++++++++++++++++++++++++++++++--------------------- 1 file changed, 48 insertions(+), 29 deletions(-) diff --git a/nodes.py b/nodes.py index 818393d..0ef571c 100644 --- a/nodes.py +++ b/nodes.py @@ -916,6 +916,7 @@ class WanVideoImageToVideoEncode: "tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}), "add_cond_latents": ("ADD_COND_LATENTS", {"advanced": True, "tooltip": "Additional cond latents WIP"}), "augment_empty_frames": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "EXPERIMENTAL: Augment empty frames with the difference to the start image to force more motion"}), + "empty_frame_pad_image": ("IMAGE", {"tooltip": "Use this image to pad empty frames instead of gray, used with SVI-shot and SVI 2.0 LoRAs"}), } } @@ -925,14 +926,14 @@ class WanVideoImageToVideoEncode: CATEGORY = "WanVideoWrapper" def process(self, width, height, num_frames, force_offload, noise_aug_strength, - start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False, - temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None, vae=None, augment_empty_frames=0.0): - + start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False, + temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, add_cond_latents=None, vae=None, augment_empty_frames=0.0, empty_frame_pad_image=None): + if vae is None: raise ValueError("VAE is required for image encoding.") H = height W = width - + lat_h = H // vae.upsampling_factor lat_w = W // vae.upsampling_factor @@ -957,6 +958,8 @@ class WanVideoImageToVideoEncode: mask = torch.cat([mask, torch.zeros(base_frames - mask.shape[0], lat_h, lat_w, device=device)]) mask = mask.unsqueeze(0).to(device, vae.dtype) + pixel_mask = mask.clone() + # Repeat first frame and optionally end frame start_mask_repeated = torch.repeat_interleave(mask[:, 0:1], repeats=4, dim=1) # T, C, H, W if end_image is not None and not fun_or_fl2v_model: @@ -979,7 +982,7 @@ class WanVideoImageToVideoEncode: resized_start_image = resized_start_image * 2 - 1 if noise_aug_strength > 0.0: resized_start_image = add_noise_to_reference_video(resized_start_image, ratio=noise_aug_strength) - + if end_image is not None: end_image = end_image[..., :3] if end_image.shape[1] != H or end_image.shape[2] != W: @@ -989,30 +992,46 @@ class WanVideoImageToVideoEncode: resized_end_image = resized_end_image * 2 - 1 if noise_aug_strength > 0.0: resized_end_image = add_noise_to_reference_video(resized_end_image, ratio=noise_aug_strength) - + # Concatenate image with zero frames and encode - if temporal_mask is None: - if start_image is not None and end_image is None: - zero_frames = torch.zeros(3, num_frames-start_image.shape[0], H, W, device=device, dtype=vae.dtype) - concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames], dim=1) - del resized_start_image, zero_frames - elif start_image is None and end_image is not None: - zero_frames = torch.zeros(3, num_frames-end_image.shape[0], H, W, device=device, dtype=vae.dtype) - concatenated = torch.cat([zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1) - del zero_frames - elif start_image is None and end_image is None: - concatenated = torch.zeros(3, num_frames, H, W, device=device, dtype=vae.dtype) - else: - if fun_or_fl2v_model: - zero_frames = torch.zeros(3, num_frames-(start_image.shape[0]+end_image.shape[0]), H, W, device=device, dtype=vae.dtype) - else: - zero_frames = torch.zeros(3, num_frames-1, H, W, device=device, dtype=vae.dtype) - concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1) - del resized_start_image, zero_frames + if start_image is not None and end_image is None: + zero_frames = torch.zeros(3, num_frames-start_image.shape[0], H, W, device=device, dtype=vae.dtype) + concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames], dim=1) + del resized_start_image, zero_frames + elif start_image is None and end_image is not None: + zero_frames = torch.zeros(3, num_frames-end_image.shape[0], H, W, device=device, dtype=vae.dtype) + concatenated = torch.cat([zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1) + del zero_frames + elif start_image is None and end_image is None: + concatenated = torch.zeros(3, num_frames, H, W, device=device, dtype=vae.dtype) else: - temporal_mask = common_upscale(temporal_mask.unsqueeze(1), W, H, "nearest", "disabled").squeeze(1) - concatenated = resized_start_image[:,:num_frames].to(vae.dtype)# * temporal_mask[:num_frames].unsqueeze(0).to(vae.dtype) - del resized_start_image, temporal_mask + if fun_or_fl2v_model: + zero_frames = torch.zeros(3, num_frames-(start_image.shape[0]+end_image.shape[0]), H, W, device=device, dtype=vae.dtype) + else: + zero_frames = torch.zeros(3, num_frames-1, H, W, device=device, dtype=vae.dtype) + concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1) + del resized_start_image, zero_frames + + if empty_frame_pad_image is not None: + pad_img = empty_frame_pad_image.clone()[..., :3] + if pad_img.shape[1] != H or pad_img.shape[2] != W: + pad_img = common_upscale(pad_img.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(1, -1) + pad_img = (pad_img.movedim(-1, 0) * 2 - 1).to(device, dtype=vae.dtype) + + num_pad_frames = pad_img.shape[1] + num_target_frames = concatenated.shape[1] + if num_pad_frames < num_target_frames: + pad_img = torch.cat([pad_img, pad_img[:, -1:].expand(-1, num_target_frames - num_pad_frames, -1, -1)], dim=1) + else: + pad_img = pad_img[:, :num_target_frames] + + frame_is_empty = (pixel_mask[0].mean(dim=(-2, -1)) < 0.5)[:concatenated.shape[1]].clone() + if start_image is not None: + frame_is_empty[:start_image.shape[0]] = False + if end_image is not None: + frame_is_empty[-end_image.shape[0]:] = False + + concatenated[:, frame_is_empty] = pad_img[:, frame_is_empty] mm.soft_empty_cache() gc.collect() @@ -1041,7 +1060,7 @@ class WanVideoImageToVideoEncode: if add_cond_latents is not None: add_cond_latents["ref_latent_neg"] = vae.encode(torch.zeros(1, 3, 1, H, W, device=device, dtype=vae.dtype), device) - + if force_offload: vae.model.to(offload_device) mm.soft_empty_cache() @@ -1064,7 +1083,7 @@ class WanVideoImageToVideoEncode: } return (image_embeds,) - + # region WanAnimate class WanVideoAnimateEmbeds: @classmethod