InfiniteTalk: Pad with silent embeds instead of repeat, add MultiTalkSilentEmbeds -node

This commit is contained in:
kijai
2025-09-06 21:58:03 +03:00
parent ae768f53a4
commit 011c0ce38d
3 changed files with 50 additions and 6 deletions
Binary file not shown.
+33 -2
View File
@@ -327,6 +327,35 @@ class MultiTalkWav2VecEmbeds:
log.info(f"[MultiTalk] multi_audio_type={multi_audio_type} | final waveform shape={out_audio['waveform'].shape} | length={out_audio['waveform'].shape[-1]} samples | seconds={out_audio['waveform'].shape[-1]/sr:.3f}s (expected {'sum' if multi_audio_type=='add' else 'max'} of raw)")
return (multitalk_embeds, out_audio, actual_num_frames)
class MultiTalkSilentEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 1, "tooltip": "The total frame count to generate."}),
},
}
RETURN_TYPES = ("MULTITALK_EMBEDS", )
RETURN_NAMES = ("multitalk_embeds", )
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, num_frames):
silence_path = os.path.join(script_directory, "encoded_silence.safetensors")
encoded_silence = load_torch_file(silence_path)["audio_emb"]
target_frames = num_frames
repeats = (target_frames + encoded_silence.shape[0] - 1) // encoded_silence.shape[0]
repeated = encoded_silence.repeat(repeats, 1, 1)
repeated = repeated[:target_frames]
multitalk_embeds = {
"audio_features": repeated,
"audio_scale": 1.0,
"audio_cfg_scale": 1.0,
"ref_target_masks": None
}
return (multitalk_embeds,)
class WanVideoImageToVideoMultiTalk:
@@ -410,12 +439,14 @@ NODE_CLASS_MAPPINGS = {
"MultiTalkModelLoader": MultiTalkModelLoader,
"MultiTalkWav2VecEmbeds": MultiTalkWav2VecEmbeds,
"WanVideoImageToVideoMultiTalk": WanVideoImageToVideoMultiTalk,
"Wav2VecModelLoader": Wav2VecModelLoader
"Wav2VecModelLoader": Wav2VecModelLoader,
"MultiTalkSilentEmbeds": MultiTalkSilentEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"MultiTalkModelLoader": "Multi/InfiniteTalk Model Loader",
"MultiTalkWav2VecEmbeds": "Multi/InfiniteTalk Wav2vec2 Embeds",
"WanVideoImageToVideoMultiTalk": "WanVideo Long I2V Multi/InfiniteTalk",
"Wav2VecModelLoader": "Wav2vec2 Model Loader"
"Wav2VecModelLoader": "Wav2vec2 Model Loader",
"MultiTalkSilentEmbeds": "MultiTalk Silent Embeds",
}
+17 -4
View File
@@ -24,7 +24,7 @@ from contextlib import nullcontext
from einops import rearrange
from comfy import model_management as mm
from comfy.utils import ProgressBar, common_upscale
from comfy.utils import ProgressBar, common_upscale, load_torch_file
from comfy.clip_vision import clip_preprocess, ClipVisionModel
from comfy.cli_args import args, LatentPreviewMethod
import folder_paths
@@ -3429,6 +3429,14 @@ class WanVideoSampler:
"end": uni3c_embeds["end"],
}
encoded_silence = None
try:
silence_path = os.path.join(script_directory, "multitalk", "encoded_silence.safetensors")
encoded_silence = load_torch_file(silence_path)["audio_emb"].to(dtype)
except:
log.warning("No encoded silence file found, padding with end of audio embedding instead.")
total_frames = len(audio_embedding[0])
estimated_iterations = total_frames // (frame_num - motion_frame) + 1
callback = prepare_callback(patcher, estimated_iterations)
@@ -3768,9 +3776,14 @@ class WanVideoSampler:
source_frame = len(audio_embedding[human_inx])
source_frames.append(source_frame)
if audio_end_idx >= len(audio_embedding[human_inx]):
miss_length = audio_end_idx - len(audio_embedding[human_inx]) + 3
add_audio_emb = torch.flip(audio_embedding[human_inx][-1*miss_length:], dims=[0])
audio_embedding[human_inx] = torch.cat([audio_embedding[human_inx], add_audio_emb], dim=0)
print(f"Audio embedding for subject {human_inx} not long enough: {len(audio_embedding[human_inx])}, need {audio_end_idx}, padding...")
miss_length = audio_end_idx - len(audio_embedding[human_inx]) + 3
print(f"Padding length: {miss_length}")
if encoded_silence is not None:
add_audio_emb = encoded_silence[-1*miss_length:]
else:
add_audio_emb = torch.flip(audio_embedding[human_inx][-1*miss_length:], dims=[0])
audio_embedding[human_inx] = torch.cat([audio_embedding[human_inx], add_audio_emb.to(device, dtype)], dim=0)
miss_lengths.append(miss_length)
else:
miss_lengths.append(0)