diff --git a/multitalk/nodes.py b/multitalk/nodes.py index f1d8d68..733a4aa 100644 --- a/multitalk/nodes.py +++ b/multitalk/nodes.py @@ -62,12 +62,26 @@ class MultiTalkModelLoader: return (multitalk,) + +def loudness_norm(audio_array, sr=16000, lufs=-23): + try: + import pyloudnorm + except: + raise ImportError("pyloudnorm package is not installed") + meter = pyloudnorm.Meter(sr) + loudness = meter.integrated_loudness(audio_array) + if abs(loudness) > 100: + return audio_array + normalized_audio = pyloudnorm.normalize.loudness(audio_array, loudness, lufs) + return normalized_audio + class MultiTalkWav2VecEmbeds: @classmethod def INPUT_TYPES(s): return {"required": { "wav2vec_model": ("WAV2VECMODEL",), "audio": ("AUDIO",), + "normalize_loudness": ("BOOLEAN", {"default": True}), "num_frames": ("INT", {"default": 81, "min": 1, "max": 1000, "step": 1}), "fps": ("FLOAT", {"default": 23.0, "min": 1.0, "max": 60.0, "step": 0.1}), "audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1, "tooltip": "Strength of the audio conditioning"}), @@ -75,12 +89,12 @@ class MultiTalkWav2VecEmbeds: }, } - RETURN_TYPES = ("MULTITALK_EMBEDS", ) - RETURN_NAMES = ("multitalk_embeds",) + RETURN_TYPES = ("MULTITALK_EMBEDS", "AUDIO", ) + RETURN_NAMES = ("multitalk_embeds", "audio", ) FUNCTION = "process" CATEGORY = "WanVideoWrapper" - def process(self, wav2vec_model, fps, num_frames, audio, audio_scale, audio_cfg_scale): + def process(self, wav2vec_model, normalize_loudness, fps, num_frames, audio, audio_scale, audio_cfg_scale): import torchaudio import numpy as np from einops import rearrange @@ -110,10 +124,13 @@ class MultiTalkWav2VecEmbeds: except: audio_segment = audio_input - print("audio_segment.shape", audio_segment.shape) + audio_segment = audio_segment.numpy() + + if normalize_loudness: + audio_segment = loudness_norm(audio_segment, sr=sr) audio_feature = np.squeeze( - wav2vec_feature_extractor(audio_segment.numpy(), sampling_rate=sr).input_values + wav2vec_feature_extractor(audio_segment, sampling_rate=sr).input_values ) audio_feature = torch.from_numpy(audio_feature).float().to(device=device) @@ -134,8 +151,13 @@ class MultiTalkWav2VecEmbeds: "audio_scale": audio_scale, "audio_cfg_scale": audio_cfg_scale } + + audio_output = { + "waveform": audio_feature.unsqueeze(0).cpu(), + "sample_rate": sr + } - return (multitalk_embeds,) + return (multitalk_embeds, audio_output) NODE_CLASS_MAPPINGS = { "MultiTalkModelLoader": MultiTalkModelLoader, diff --git a/nodes.py b/nodes.py index 1ab9f4e..47d1132 100644 --- a/nodes.py +++ b/nodes.py @@ -3291,11 +3291,10 @@ class WanVideoSampler: ).unsqueeze( 1 ) + indices.unsqueeze(0) - center_indices = torch.clamp(center_indices, min=0, max=audio_embedding[human_idx].shape[0]-1) - audio_emb = audio_embedding[human_idx][center_indices][None,...].to(device) + center_indices = torch.clamp(center_indices, min=0, max=audio_embedding[human_idx].shape[0] - 1) + audio_emb = audio_embedding[human_idx][center_indices].unsqueeze(0).to(device) audio_embs.append(audio_emb) audio_embs = torch.concat(audio_embs, dim=0).to(dtype) - print("audio_embs: ", audio_embs.shape) base_params = { 'seq_len': seq_len,