You can integrate nnAudio layers directly into a torch.nn.Module. This allows your model to accept raw waveforms as input and perform spectrogram extraction automatically during the forward pass. This is useful for end-to-end trainable audio models.
Example of a model that performs on-the-fly STFT extraction followed by CNN layers:
from nnAudio import features
import torch
import torch.nn as nn
class Model(torch.nn.Module):
def __init__(self, n_fft, output_dim):
super().__init__()
self.epsilon = 1e-10
# Getting Mel Spectrogram on the fly
self.spec_layer = features.STFT(n_fft=n_fft, freq_bins=None,
hop_length=512, window='hann',
freq_scale='no', center=True,
pad_mode='reflect', fmin=50,
fmax=6000, sr=22050, trainable=False,
output_format='Magnitude')
self.n_bins = n_fft // 2
# Creating CNN Layers
self.CNN_freq_kernel_size = (128, 1)
self.CNN_freq_kernel_stride = (2, 1)
k_out = 128
k2_out = 256
self.CNN_freq = nn.Conv2d(1, k_out,
kernel_size=self.CNN_freq_kernel_size,
stride=self.CNN_freq_kernel_stride)
self.CNN_time = nn.Conv2d(k_out, k2_out,
kernel_size=(1, 3), stride=(1, 1))
self.region_v = 1 + (self.n_bins - self.CNN_freq_kernel_size[0]) // self.CNN_freq_kernel_stride[0]
self.linear = torch.nn.Linear(k2_out * self.region_v, output_dim, bias=False)
def forward(self, x):
z = self.spec_layer(x)
z = torch.log(z + self.epsilon)
z2 = torch.relu(self.CNN_freq(z.unsqueeze(1)))
z3 = torch.relu(self.CNN_time(z2)).mean(-1)
y = self.linear(torch.relu(torch.flatten(z3, 1)))
return torch.sigmoid(y)
# Usage: model takes waveforms directly
model = Model(n_fft=1024, output_dim=10)
waveforms = torch.randn(4, 44100)
output = model(waveforms) # automatically converts waveforms into spectrograms