InfiniteTalk: Add option to save results during the process, make it clearer no latents are returned

This commit is contained in:
kijai
2025-09-08 18:50:53 +03:00
parent e5955d8395
commit fe379fa77f
4 changed files with 496 additions and 479 deletions
+15 -6
View File
@@ -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,
+15 -1
View File
@@ -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
View File
@@ -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",
}