InfiniteTalk: Pad with silent embeds instead of repeat, add MultiTalkSilentEmbeds -node
This commit is contained in:
Binary file not shown.
+33
-2
@@ -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",
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user