import torch import torchaudio from torch import nn from librosa import filters from nnAudio.features import CQT1992v2, CQT2010v2 def get_chroma_filterbank(sample_rate, n_fft): return torch.tensor(filters.chroma(sr=sample_rate, n_fft=n_fft).astype("float32")) class STFT(nn.Module): def __init__( self, n_fft=2048, hop_length=240, is_db=False, ): super(STFT, self).__init__() # short-time Fourier transform self.stft = torchaudio.transforms.Spectrogram( n_fft=n_fft, hop_length=hop_length ) # amplitude to decibel self.is_db = is_db if is_db: self.amplitude_to_db = torchaudio.transforms.AmplitudeToDB() def forward(self, waveform): if self.is_db: return self.amplitude_to_db(self.stft(waveform)) else: return self.stft(waveform) class MultiResSTFT(nn.Module): def __init__( self, n_ffts=[256, 512, 1024, 2048, 4096], hop_length=240, is_db=False, ): super(MultiResSTFT, self).__init__() # multi-resolution short-time Fourier transform self.n_ffts = n_ffts for n_fft in n_ffts: stft = torchaudio.transforms.Spectrogram(n_fft=n_fft, hop_length=hop_length) setattr(self, "stft_%d" % n_fft, stft) # amplitude to decibel self.is_db = is_db if is_db: self.amplitude_to_db = torchaudio.transforms.AmplitudeToDB() def forward(self, waveform): specs = [] for n_fft in self.n_ffts: stft = getattr(self, "stft_%d" % n_fft) spec = stft(waveform) if self.is_db: spec = self.amplitude_to_db(spec) specs.append(spec) return specs class MelSTFT(nn.Module): def __init__( self, sample_rate=24000, n_fft=2048, hop_length=240, n_mels=128, is_db=False, ): super(MelSTFT, self).__init__() # short-time Fourier transform with mel filterbank self.mel_stft = torchaudio.transforms.MelSpectrogram( sample_rate=sample_rate, n_fft=n_fft, hop_length=hop_length, n_mels=n_mels, ) # amplitude to decibel self.is_db = is_db if is_db: self.amplitude_to_db = torchaudio.transforms.AmplitudeToDB() def forward(self, waveform): if self.is_db: return self.amplitude_to_db(self.mel_stft(waveform)) else: return self.mel_stft(waveform) class MFCC(nn.Module): def __init__( self, sample_rate=24000, n_fft=2048, hop_length=240, n_mels=128, n_mfcc=13, ): super(MFCC, self).__init__() # MFCC self.mfcc = torchaudio.transforms.MFCC( sample_rate=sample_rate, n_mfcc=n_mfcc, melkwargs={ "n_fft": n_fft, "hop_length": hop_length, "n_mels": n_mels, }, ) def forward(self, waveform): return self.mfcc(waveform)[:, 1:, :] # first channel is energy class Chromagram(nn.Module): def __init__( self, sample_rate=24000, n_fft=2048, hop_length=240, ): super(Chromagram, self).__init__() # short-time Fourier transform self.stft = torchaudio.transforms.Spectrogram( n_fft=n_fft, hop_length=hop_length, ) # chroma filterbank self.register_buffer("chroma_fb", get_chroma_filterbank(sample_rate, n_fft)) def forward(self, waveform): spec = self.stft(waveform) chromagram = torch.matmul(spec.transpose(1, 2), self.chroma_fb.T).transpose( 1, 2 ) return chromagram class CQT(nn.Module): def __init__( self, sample_rate=24000, hop_length=240, algorithm="1992", is_db=False ): super().__init__() assert algorithm in ["1992", "2010"] # constant-Q transform if algorithm == "1992": self.cqt = CQT1992v2(sr=sample_rate, hop_length=hop_length) elif algorithm == "2010": self.cqt = CQT2010v2(sr=sample_rate, hop_length=hop_length) # amplitude to decibel self.is_db = is_db if is_db: self.amplitude_to_db = torchaudio.transforms.AmplitudeToDB() def forward(self, waveform): if self.is_db: return self.amplitude_to_db(self.cqt(waveform)) else: return self.cqt(waveform) """ TODO: Band-split spec """