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
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user