diff --git a/multitalk/encoded_silence.safetensors b/multitalk/encoded_silence.safetensors new file mode 100644 index 0000000..5ab6aa7 Binary files /dev/null and b/multitalk/encoded_silence.safetensors differ diff --git a/multitalk/nodes.py b/multitalk/nodes.py index 2f30f53..dd2da93 100644 --- a/multitalk/nodes.py +++ b/multitalk/nodes.py @@ -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", } \ No newline at end of file diff --git a/nodes.py b/nodes.py index 4255168..507e5b8 100644 --- a/nodes.py +++ b/nodes.py @@ -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)