Views
No views yet
1class MelSpectrogramFeatures(FeatureExtractor):
2 def __init__(self, sample_rate=24000, n_fft=1024, hop_length=256, n_mels=100, padding="center"):
3 super().__init__()
4 if padding not in ["center", "same"]:
5 raise ValueError("Padding must be 'center' or 'same'.")
6 self.padding = padding
7 self.mel_spec = torchaudio.transforms.MelSpectrogram(
8 sample_rate=sample_rate,
9 n_fft=n_fft,
10 hop_length=hop_length,
11 n_mels=n_mels,
12 center=padding == "center",
13 power=1,
14 )
15
16 def forward(self, audio, **kwargs):
17 if self.padding == "same":
18 pad = self.mel_spec.win_length - self.mel_spec.hop_length
19 audio = torch.nn.functional.pad(audio, (pad // 2, pad // 2), mode="reflect")
20 mel = self.mel_spec(audio)
21 features = safe_log(mel)
22 return features