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"}),
|
"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"}),
|
"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"}),
|
"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"
|
CATEGORY = "WanVideoWrapper"
|
||||||
|
|
||||||
def process(self, width, height, num_frames, force_offload, noise_aug_strength,
|
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,
|
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):
|
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:
|
if vae is None:
|
||||||
raise ValueError("VAE is required for image encoding.")
|
raise ValueError("VAE is required for image encoding.")
|
||||||
H = height
|
H = height
|
||||||
W = width
|
W = width
|
||||||
|
|
||||||
lat_h = H // vae.upsampling_factor
|
lat_h = H // vae.upsampling_factor
|
||||||
lat_w = W // 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 = torch.cat([mask, torch.zeros(base_frames - mask.shape[0], lat_h, lat_w, device=device)])
|
||||||
mask = mask.unsqueeze(0).to(device, vae.dtype)
|
mask = mask.unsqueeze(0).to(device, vae.dtype)
|
||||||
|
|
||||||
|
pixel_mask = mask.clone()
|
||||||
|
|
||||||
# Repeat first frame and optionally end frame
|
# Repeat first frame and optionally end frame
|
||||||
start_mask_repeated = torch.repeat_interleave(mask[:, 0:1], repeats=4, dim=1) # T, C, H, W
|
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:
|
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
|
resized_start_image = resized_start_image * 2 - 1
|
||||||
if noise_aug_strength > 0.0:
|
if noise_aug_strength > 0.0:
|
||||||
resized_start_image = add_noise_to_reference_video(resized_start_image, ratio=noise_aug_strength)
|
resized_start_image = add_noise_to_reference_video(resized_start_image, ratio=noise_aug_strength)
|
||||||
|
|
||||||
if end_image is not None:
|
if end_image is not None:
|
||||||
end_image = end_image[..., :3]
|
end_image = end_image[..., :3]
|
||||||
if end_image.shape[1] != H or end_image.shape[2] != W:
|
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
|
resized_end_image = resized_end_image * 2 - 1
|
||||||
if noise_aug_strength > 0.0:
|
if noise_aug_strength > 0.0:
|
||||||
resized_end_image = add_noise_to_reference_video(resized_end_image, ratio=noise_aug_strength)
|
resized_end_image = add_noise_to_reference_video(resized_end_image, ratio=noise_aug_strength)
|
||||||
|
|
||||||
# Concatenate image with zero frames and encode
|
# Concatenate image with zero frames and encode
|
||||||
if temporal_mask is None:
|
if start_image is not None and end_image 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)
|
||||||
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)
|
||||||
concatenated = torch.cat([resized_start_image.to(device, dtype=vae.dtype), zero_frames], dim=1)
|
del resized_start_image, zero_frames
|
||||||
del resized_start_image, zero_frames
|
elif start_image is None and end_image is not None:
|
||||||
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)
|
||||||
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)
|
||||||
concatenated = torch.cat([zero_frames, resized_end_image.to(device, dtype=vae.dtype)], dim=1)
|
del zero_frames
|
||||||
del zero_frames
|
elif start_image is None and end_image is None:
|
||||||
elif start_image is None and end_image is None:
|
concatenated = torch.zeros(3, num_frames, H, W, device=device, dtype=vae.dtype)
|
||||||
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
|
|
||||||
else:
|
else:
|
||||||
temporal_mask = common_upscale(temporal_mask.unsqueeze(1), W, H, "nearest", "disabled").squeeze(1)
|
if fun_or_fl2v_model:
|
||||||
concatenated = resized_start_image[:,:num_frames].to(vae.dtype)# * temporal_mask[:num_frames].unsqueeze(0).to(vae.dtype)
|
zero_frames = torch.zeros(3, num_frames-(start_image.shape[0]+end_image.shape[0]), H, W, device=device, dtype=vae.dtype)
|
||||||
del resized_start_image, temporal_mask
|
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()
|
mm.soft_empty_cache()
|
||||||
gc.collect()
|
gc.collect()
|
||||||
@@ -1041,7 +1060,7 @@ class WanVideoImageToVideoEncode:
|
|||||||
|
|
||||||
if add_cond_latents is not None:
|
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)
|
add_cond_latents["ref_latent_neg"] = vae.encode(torch.zeros(1, 3, 1, H, W, device=device, dtype=vae.dtype), device)
|
||||||
|
|
||||||
if force_offload:
|
if force_offload:
|
||||||
vae.model.to(offload_device)
|
vae.model.to(offload_device)
|
||||||
mm.soft_empty_cache()
|
mm.soft_empty_cache()
|
||||||
@@ -1064,7 +1083,7 @@ class WanVideoImageToVideoEncode:
|
|||||||
}
|
}
|
||||||
|
|
||||||
return (image_embeds,)
|
return (image_embeds,)
|
||||||
|
|
||||||
# region WanAnimate
|
# region WanAnimate
|
||||||
class WanVideoAnimateEmbeds:
|
class WanVideoAnimateEmbeds:
|
||||||
@classmethod
|
@classmethod
|
||||||
|
|||||||
Reference in New Issue
Block a user