InfiniteTalk: Add option to save results during the process, make it clearer no latents are returned
This commit is contained in:
+442
-470
File diff suppressed because it is too large
Load Diff
+15
-6
@@ -6,6 +6,7 @@ import torch
|
||||
from ..utils import log, set_module_tensor_to_device
|
||||
import os
|
||||
import json
|
||||
import datetime
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
folder_paths.add_model_folder_path("wav2vec2", os.path.join(folder_paths.models_dir, "wav2vec2"))
|
||||
@@ -389,17 +390,19 @@ class WanVideoImageToVideoMultiTalk:
|
||||
"auto",
|
||||
"multitalk",
|
||||
"infinitetalk"
|
||||
], {"default": "auto", "tooltip": "The sampling strategy to use in the long video generation loop, should match the model used"})
|
||||
], {"default": "auto", "tooltip": "The sampling strategy to use in the long video generation loop, should match the model used"}),
|
||||
"output_path": ("STRING", {"default": "", "tooltip": "If set, will save each window's resulting frames to this folder"}),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "STRING",)
|
||||
RETURN_NAMES = ("image_embeds", "output_path")
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Enables Multi/InfiniteTalk long video generation sampling method, the video is created in windows with overlapping frames. Not compatible or necessary to be used with context windows and many other features besides Multi/InfiniteTalk."
|
||||
|
||||
def process(self, vae, width, height, frame_window_size, motion_frame, force_offload, colormatch, start_image=None, tiled_vae=False, clip_embeds=None, mode="multitalk"):
|
||||
def process(self, vae, width, height, frame_window_size, motion_frame, force_offload, colormatch, start_image=None, tiled_vae=False, clip_embeds=None, mode="multitalk", output_path=""):
|
||||
|
||||
H = height
|
||||
W = width
|
||||
@@ -416,6 +419,11 @@ class WanVideoImageToVideoMultiTalk:
|
||||
target_shape = (16, (num_frames - 1) // VAE_STRIDE[0] + 1,
|
||||
height // VAE_STRIDE[1],
|
||||
width // VAE_STRIDE[2])
|
||||
|
||||
if output_path:
|
||||
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
output_path = os.path.join(output_path, f"{timestamp}_{mode}_output")
|
||||
os.makedirs(output_path, exist_ok=True)
|
||||
|
||||
image_embeds = {
|
||||
"multitalk_sampling": True,
|
||||
@@ -430,10 +438,11 @@ class WanVideoImageToVideoMultiTalk:
|
||||
"target_shape": target_shape,
|
||||
"clip_context": clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None,
|
||||
"colormatch": colormatch,
|
||||
"multitalk_mode": mode
|
||||
"multitalk_mode": mode,
|
||||
"output_path": output_path
|
||||
}
|
||||
|
||||
return (image_embeds,)
|
||||
return (image_embeds, output_path)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"MultiTalkModelLoader": MultiTalkModelLoader,
|
||||
|
||||
@@ -5,6 +5,7 @@ import numpy as np
|
||||
from tqdm import tqdm
|
||||
import inspect
|
||||
import copy
|
||||
from PIL import Image
|
||||
import hashlib
|
||||
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
|
||||
|
||||
@@ -3391,6 +3392,9 @@ class WanVideoSampler:
|
||||
if original_images is None:
|
||||
original_images = torch.zeros([noise.shape[0], 1, target_h, target_w], device=device)
|
||||
|
||||
output_path = image_embeds.get("output_path", "")
|
||||
img_counter = 0
|
||||
|
||||
if len(multitalk_embeds['audio_features'])==2 and (multitalk_embeds['ref_target_masks'] is None):
|
||||
face_scale = 0.1
|
||||
x_min, x_max = int(target_h * face_scale), int(target_h * (1 - face_scale))
|
||||
@@ -3746,7 +3750,17 @@ class WanVideoSampler:
|
||||
|
||||
videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2)
|
||||
|
||||
# cache generated samples
|
||||
# optionally save generated samples to disk
|
||||
if output_path:
|
||||
video_np = videos.clamp(-1.0, 1.0).add(1.0).div(2.0).mul(255).cpu().float().numpy().transpose(1, 2, 3, 0).astype('uint8')
|
||||
num_frames_to_save = video_np.shape[0] if is_first_clip else video_np.shape[0] - cur_motion_frames_num
|
||||
log.info(f"Saving {num_frames_to_save} generated frames to {output_path}")
|
||||
start_idx = 0 if is_first_clip else cur_motion_frames_num
|
||||
for i in range(start_idx, video_np.shape[0]):
|
||||
im = Image.fromarray(video_np[i])
|
||||
im.save(os.path.join(output_path, f"frame_{img_counter:05d}.png"))
|
||||
img_counter += 1
|
||||
|
||||
gen_video_list.append(videos if is_first_clip else videos[:, cur_motion_frames_num:])
|
||||
|
||||
current_condframe_index += 1
|
||||
|
||||
+24
-2
@@ -448,6 +448,26 @@ class NormalizeAudioLoudness:
|
||||
normalized_audio = pyloudnorm.normalize.loudness(audio_array, loudness, lufs)
|
||||
return normalized_audio
|
||||
|
||||
class WanVideoPassImagesFromSamples:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"samples": ("LATENT",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("images",)
|
||||
FUNCTION = "decode"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Gets possible already decoded images from the samples dictionary, used with Multi/InfiniteTalk sampling"
|
||||
|
||||
def decode(self, samples):
|
||||
video = samples.get("video", None)
|
||||
video.clamp_(-1.0, 1.0)
|
||||
video.add_(1.0).div_(2.0)
|
||||
return video.cpu().float(),
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoImageResizeToClosest": WanVideoImageResizeToClosest,
|
||||
"WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame,
|
||||
@@ -457,7 +477,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoLatentReScale": WanVideoLatentReScale,
|
||||
"CreateScheduleFloatList": CreateScheduleFloatList,
|
||||
"WanVideoSigmaToStep": WanVideoSigmaToStep,
|
||||
"NormalizeAudioLoudness": NormalizeAudioLoudness
|
||||
"NormalizeAudioLoudness": NormalizeAudioLoudness,
|
||||
"WanVideoPassImagesFromSamples": WanVideoPassImagesFromSamples,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoImageResizeToClosest": "WanVideo Image Resize To Closest",
|
||||
@@ -468,5 +489,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoLatentReScale": "WanVideo Latent ReScale",
|
||||
"CreateScheduleFloatList": "Create Schedule Float List",
|
||||
"WanVideoSigmaToStep": "WanVideo Sigma To Step",
|
||||
"NormalizeAudioLoudness": "Normalize Audio Loudness"
|
||||
"NormalizeAudioLoudness": "Normalize Audio Loudness",
|
||||
"WanVideoPassImagesFromSamples": "WanVideo Pass Images From Samples",
|
||||
}
|
||||
Reference in New Issue
Block a user