From 8d0fcee2f3decc1cbbf3b850332e6b2a022e1377 Mon Sep 17 00:00:00 2001 From: Mel Massadian Date: Sun, 28 Jul 2024 20:11:31 +0200 Subject: [PATCH] =?UTF-8?q?feat:=20=E2=9C=A8=20add=20AudioStack?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit To stack/overlay audios. --- nodes/audio.py | 155 +++++++++++++++++++++++++++++++++++---------- web/mtb_widgets.js | 3 +- 2 files changed, 125 insertions(+), 33 deletions(-) diff --git a/nodes/audio.py b/nodes/audio.py index 1f7a925..1224a1d 100644 --- a/nodes/audio.py +++ b/nodes/audio.py @@ -1,8 +1,121 @@ +from typing import TypedDict + import torch import torchaudio -class MTB_AudioSequence: +class AudioDict(TypedDict): + """Comfy's representation of AUDIO data.""" + + sample_rate: int + waveform: torch.Tensor + + +AudioData = AudioDict | list[AudioDict] + + +class MtbAudio: + """Base class for audio processing.""" + + @classmethod + def is_stereo( + cls, + audios: AudioData, + ) -> bool: + if isinstance(audios, list): + return any(cls.is_stereo(audio) for audio in audios) + else: + return audios["waveform"].shape[1] == 2 + + @staticmethod + def resample(audio: AudioDict, common_sample_rate: int) -> AudioDict: + if audio["sample_rate"] != common_sample_rate: + resampler = torchaudio.transforms.Resample( + orig_freq=audio["sample_rate"], new_freq=common_sample_rate + ) + return { + "sample_rate": common_sample_rate, + "waveform": resampler(audio["waveform"]), + } + else: + return audio + + @staticmethod + def to_stereo(audio: AudioDict) -> AudioDict: + if audio["waveform"].shape[1] == 1: + return { + "sample_rate": audio["sample_rate"], + "waveform": torch.cat( + [audio["waveform"], audio["waveform"]], dim=1 + ), + } + else: + return audio + + @classmethod + def preprocess_audios( + cls, audios: list[AudioDict] + ) -> tuple[list[AudioDict], bool, int]: + max_sample_rate = max([audio["sample_rate"] for audio in audios]) + + resampled_audios = [ + cls.resample(audio, max_sample_rate) for audio in audios + ] + + is_stereo = cls.is_stereo(audios) + if is_stereo: + audios = [cls.to_stereo(audio) for audio in resampled_audios] + + return (audios, is_stereo, max_sample_rate) + + +class MTB_AudioStack(MtbAudio): + """Stack/Overlay audio inputs (dynamic inputs). + + - pad audios to the longest inputs. + - resample audios to the highest sample rate in the inputs. + - convert them all to stereo if one of the inputs is. + """ + + @classmethod + def INPUT_TYPES(cls): + return {"required": {}} + + RETURN_TYPES = ("AUDIO",) + RETURN_NAMES = ("stacked_audio",) + CATEGORY = "mtb/audio" + FUNCTION = "stack" + + def stack(self, **kwargs: AudioDict) -> tuple[AudioDict]: + audios, is_stereo, max_rate = self.preprocess_audios( + list(kwargs.values()) + ) + + max_length = max([audio["waveform"].shape[-1] for audio in audios]) + + padded_audios: list[torch.Tensor] = [] + for audio in audios: + padding = torch.zeros( + ( + 1, + 2 if is_stereo else 1, + max_length - audio["waveform"].shape[-1], + ) + ) + padded_audio = torch.cat([audio["waveform"], padding], dim=-1) + padded_audios.append(padded_audio) + + stacked_waveform = torch.stack(padded_audios, dim=0).sum(dim=0) + + return ( + { + "sample_rate": max_rate, + "waveform": stacked_waveform, + }, + ) + + +class MTB_AudioSequence(MtbAudio): """Sequence audio inputs (dynamic inputs). - adding silence_duration between each segment. @@ -23,53 +136,31 @@ class MTB_AudioSequence: CATEGORY = "mtb/audio" FUNCTION = "sequence" - def sequence(self, silence_duration: float, **kwargs): - audios = kwargs.values() - - common_sample_rate = max([audio["sample_rate"] for audio in audios]) - - is_stereo = any( - waveform.shape[1] == 2 - for waveform in [audio["waveform"] for audio in audios] + def sequence(self, silence_duration: float, **kwargs: AudioDict): + audios, is_stereo, max_rate = self.preprocess_audios( + list(kwargs.values()) ) - resampled_audios = [] - for audio in audios: - if audio["sample_rate"] != common_sample_rate: - resampler = torchaudio.transforms.Resample( - orig_freq=audio["sample_rate"], new_freq=common_sample_rate - ) - audio["waveform"] = resampler(audio["waveform"]) - - # convert to stereo - if is_stereo and audio["waveform"].shape[1] == 1: - audio["waveform"] = torch.cat( - [audio["waveform"], audio["waveform"]], dim=1 - ) - resampled_audios.append(audio) - silence = torch.zeros( ( 1, 2 if is_stereo else 1, - int(silence_duration * common_sample_rate), + int(silence_duration * max_rate), ) ) - sequence = [] - for i, audio in enumerate(resampled_audios): + sequence: list[torch.Tensor] = [] + for i, audio in enumerate(audios): sequence.append(audio["waveform"]) - if i < len(resampled_audios) - 1: + if i < len(audios) - 1: sequence.append(silence) sequenced_waveform = torch.cat(sequence, dim=-1) return ( { - "sample_rate": common_sample_rate, + "sample_rate": max_rate, "waveform": sequenced_waveform, }, ) -__nodes__ = [ - MTB_AudioSequence, -] +__nodes__ = [MTB_AudioSequence, MTB_AudioStack] diff --git a/web/mtb_widgets.js b/web/mtb_widgets.js index f8cc91d..4297049 100644 --- a/web/mtb_widgets.js +++ b/web/mtb_widgets.js @@ -1136,7 +1136,8 @@ const mtb_widgets = { shared.setupDynamicConnections(nodeType, 'image', 'IMAGE') break } - case 'Audio Sequence (mtb)': { + case 'Audio Sequence (mtb)': + case 'Audio Stack (mtb)': { shared.setupDynamicConnections(nodeType, 'audio', 'AUDIO') break }