Spaces:
Running on Zero
Running on Zero
| #!/usr/bin/env python3 | |
| # -*- coding: utf-8 -*- | |
| """ | |
| Created on Fri Apr 26 15:36:16 2024 | |
| @author: louis | |
| """ | |
| from torchaudio.transforms import Spectrogram as OriginalSpectrogram, InverseSpectrogram as OriginalInverseSpectrogram | |
| import torch | |
| from model.utils.tensor_ops import zero_pad | |
| default_stft_parameters = dict( | |
| n_fft=512, | |
| hop_length=256, | |
| win_length=512, | |
| window_fn=torch.hann_window, | |
| center=True, | |
| ) | |
| class Spectrogram(OriginalSpectrogram): | |
| def forward(self, waveform): | |
| if not self.center: | |
| raise NotImplementedError() | |
| waveform_padded = torch.nn.functional.pad(waveform, (self.n_fft // 2, self.n_fft // 2)) | |
| X = super().forward(waveform_padded) | |
| return X[..., 1:-1] | |
| class InverseSpectrogramCOLA(OriginalInverseSpectrogram): | |
| def forward(self, spectrogram, length=None): | |
| if not self.center: | |
| raise NotImplementedError() | |
| # pack batch as in original | |
| # spectrogram = torch.nn.functional.pad(spectrogram, (0, 1)) | |
| shape = spectrogram.size() | |
| spectrogram = spectrogram.reshape(-1, shape[-2], shape[-1]) | |
| expected_waveform_length = self.n_fft + self.hop_length * (shape[-1] - 1) | |
| c = torch.fft.irfft(spectrogram, dim=-2) | |
| waveform = torch.nn.functional.fold( | |
| c, | |
| output_size=(1, expected_waveform_length), | |
| kernel_size=(1, self.n_fft), | |
| dilation=1, | |
| padding=0, | |
| stride=(1, self.hop_length), | |
| ) | |
| waveform = waveform[..., self.n_fft // 2 :].squeeze(-3, -2) | |
| if length is not None: | |
| waveform = zero_pad(waveform, length) | |
| # unpack batch | |
| waveform = waveform.reshape(shape[:-2] + waveform.shape[-1:]) | |
| return waveform | |
| default_stft_module = Spectrogram(**default_stft_parameters, power=None) | |
| default_istft_module = InverseSpectrogramCOLA(**default_stft_parameters) | |