Return exact audio frame count if num_frames above it

This commit is contained in:
kijai
2025-08-20 21:08:25 +03:00
parent 68ef3ac468
commit 70bced06f4
3 changed files with 29 additions and 7 deletions
+2 -2
View File
@@ -20,8 +20,8 @@ class DownloadAndLoadWav2VecModel:
"required": {
"model": (
[
"facebook/wav2vec2-base-960h",
"TencentGameMate/chinese-wav2vec2-base"
"TencentGameMate/chinese-wav2vec2-base",
"facebook/wav2vec2-base-960h"
],
),
+21 -3
View File
@@ -91,8 +91,8 @@ class MultiTalkWav2VecEmbeds:
}
}
RETURN_TYPES = ("MULTITALK_EMBEDS", "AUDIO", )
RETURN_NAMES = ("multitalk_embeds", "audio", )
RETURN_TYPES = ("MULTITALK_EMBEDS", "AUDIO", "INT", )
RETURN_NAMES = ("multitalk_embeds", "audio", "num_frames", )
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
@@ -231,12 +231,30 @@ class MultiTalkWav2VecEmbeds:
offset += w.shape[-1]
out_audio = {"waveform": mixed, "sample_rate": sr}
# Calculate actual frames based on audio duration
actual_num_frames = num_frames
if len(audio_outputs) > 0:
if multi_audio_type == "para":
# For parallel mode, use the longest audio duration
max_audio_duration = max([ao["waveform"].shape[-1] / sr for ao in audio_outputs])
actual_frames_from_audio = int(max_audio_duration * fps)
else: # "add"
# For sequential mode, use the total audio duration
total_audio_duration = sum([ao["waveform"].shape[-1] / sr for ao in audio_outputs])
actual_frames_from_audio = int(total_audio_duration * fps)
# Use the smaller of requested frames or actual audio frames
actual_num_frames = min(num_frames, actual_frames_from_audio)
if actual_frames_from_audio < num_frames:
log.info(f"[MultiTalk] Audio duration ({actual_frames_from_audio} frames) is shorter than requested ({num_frames} frames). Using {actual_num_frames} frames.")
# Debug: log final mixed audio length and mode
total_samples_raw = sum([ao["waveform"].shape[-1] for ao in audio_outputs])
log.info(f"[MultiTalk] total raw duration = {total_samples_raw/sr:.3f}s")
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)
return (multitalk_embeds, out_audio, actual_num_frames)
class WanVideoImageToVideoMultiTalk:
+6 -2
View File
@@ -2616,8 +2616,9 @@ class WanVideoSampler:
from .latent_preview import prepare_callback #custom for tiny VAE previews
callback = prepare_callback(patcher, len(timesteps))
log.info(f"Input sequence length: {seq_len}")
log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
if not multitalk_sampling:
log.info(f"Input sequence length: {seq_len}")
log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
intermediate_device = device
@@ -3077,6 +3078,9 @@ class WanVideoSampler:
audio_embedding = multitalk_audio_embedding
human_num = len(audio_embedding)
audio_embs = None
log.info(f"Sampling {total_frames} frames in {estimated_iterations} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
while True: # start video generation iteratively
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
if mode == "infinitetalk":