feat: ✨ add AudioStack

To stack/overlay audios.
This commit is contained in:
Mel Massadian
2024-07-28 20:11:31 +02:00
parent 1078fc6f0f
commit 8d0fcee2f3
2 changed files with 125 additions and 33 deletions
+123 -32
View File
@@ -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]
+2 -1
View File
@@ -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
}