loudness norm

This commit is contained in:
kijai
2025-06-18 18:35:06 +03:00
parent 58104b620f
commit 1fe72d27aa
2 changed files with 30 additions and 9 deletions
+28 -6
View File
@@ -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,
+2 -3
View File
@@ -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,