Spaces:
Sleeping
Sleeping
| import math | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| from torchaudio.transforms import Resample | |
| import pyworld as pw | |
| from ....utils.audio.pitch_utils import interp_f0, resample_align_curve | |
| from .constants import * | |
| from .model import E2E0 | |
| from .spec import MelSpectrogram | |
| from .utils import to_local_average_f0, to_viterbi_f0 | |
| class RMVPE: | |
| def __init__(self, model_path, hop_length=160, device=None): | |
| self.resample_kernel = {} | |
| if device is None: | |
| self.device = 'cuda' if torch.cuda.is_available() else 'cpu' | |
| else: | |
| self.device = device | |
| self.model = E2E0(4, 1, (2, 2)).eval().to(self.device) | |
| ckpt = torch.load(model_path, map_location=self.device) | |
| self.model.load_state_dict(ckpt['model'], strict=False) | |
| self.mel_extractor = MelSpectrogram( | |
| N_MELS, SAMPLE_RATE, WINDOW_LENGTH, hop_length, None, MEL_FMIN, MEL_FMAX | |
| ).to(self.device) | |
| self.hop_length = hop_length | |
| def mel2hidden(self, mel): | |
| n_frames = mel.shape[-1] | |
| mel = F.pad(mel, (0, 32 * ((n_frames - 1) // 32 + 1) - n_frames), mode='constant') | |
| hidden = self.model(mel) | |
| return hidden[:, :n_frames] | |
| def decode(self, hidden, thred=0.03, use_viterbi=False): | |
| if use_viterbi: | |
| f0 = to_viterbi_f0(hidden, thred=thred) | |
| else: | |
| f0 = to_local_average_f0(hidden, thred=thred) | |
| return f0 | |
| def postprocess(self, f0, fmin=50, fmax=1000, audio=None, min_gap=2): | |
| if audio is not None: | |
| # this doesn't work. deprecated | |
| t = np.arange(0, f0.shape[0] * self.hop_length / 16000, self.hop_length / 16000) | |
| f0 = pw.stonemask(audio.astype(np.float64), f0.astype(np.float64), t, 16000).astype(float) | |
| f0[f0 < fmin] = 0 | |
| f0[f0 > fmax] = 0 | |
| # eliminate glitch | |
| # min_gap: if successive positive f0 positions < min_gap, zero these positions | |
| # eg: if min_gap=2, [0, 500, 500, 0] => [0, 0, 0, 0] | |
| for idx in range(f0.shape[0] - min_gap - 1): | |
| if f0[idx] == 0 and f0[idx + min_gap + 1] == 0 and np.sum(f0[idx: idx + min_gap + 2]) > 0: | |
| f0[idx: idx + min_gap + 2] = 0 | |
| return f0 | |
| def infer_from_audio(self, audio, sample_rate=16000, thred=0.03, use_viterbi=False): | |
| audio = torch.from_numpy(audio).float().unsqueeze(0).to(self.device) | |
| if sample_rate == 16000: | |
| audio_res = audio | |
| else: | |
| key_str = str(sample_rate) | |
| if key_str not in self.resample_kernel: | |
| self.resample_kernel[key_str] = Resample(sample_rate, 16000, lowpass_filter_width=128) | |
| self.resample_kernel[key_str] = self.resample_kernel[key_str].to(self.device) | |
| audio_res = self.resample_kernel[key_str](audio) | |
| mel = self.mel_extractor(audio_res, center=True) | |
| hidden = self.mel2hidden(mel) | |
| f0 = self.decode(hidden, thred=thred, use_viterbi=use_viterbi).squeeze(0) | |
| return f0 | |
| def get_pitch(self, waveform, sample_rate, hop_size, length, interp_uv=False, fmin=50, fmax=1000): | |
| f0 = self.infer_from_audio(waveform, sample_rate=sample_rate) | |
| f0 = self.postprocess(f0, fmin, fmax) | |
| uv = f0 == 0 | |
| time_step = hop_size / sample_rate | |
| f0_res = resample_align_curve(f0, 0.01, time_step, length) | |
| uv_res = resample_align_curve(uv.astype(np.float32), 0.01, time_step, length) > 0.5 | |
| if not interp_uv: | |
| f0_res[uv_res] = 0 | |
| return f0_res, uv_res | |
| def infer_from_audio_batch(self, audios, sample_rate=16000, thred=0.03, use_viterbi=False): | |
| from ....utils.commons.dataset_utils import collate_1d_or_2d | |
| if isinstance(audios, list): | |
| audios = [torch.from_numpy(audio).float() for audio in audios] | |
| sizes = [math.ceil((audio.shape[0] + 1) / self.hop_length) for audio in audios] | |
| audios = collate_1d_or_2d(audios, 0.0).to(self.device) | |
| elif isinstance(audios, torch.Tensor): | |
| sizes = None | |
| if audios.device != self.device: | |
| audios = audios.to(self.device) | |
| else: | |
| raise NotImplementedError | |
| if sample_rate == 16000: | |
| audios_res = audios | |
| else: | |
| key_str = str(sample_rate) | |
| if key_str not in self.resample_kernel: | |
| self.resample_kernel[key_str] = Resample(sample_rate, 16000, lowpass_filter_width=128) | |
| self.resample_kernel[key_str] = self.resample_kernel[key_str].to(self.device) | |
| audios_res = self.resample_kernel[key_str](audios) | |
| mels = self.mel_extractor(audios_res, center=True) | |
| hiddens = self.mel2hidden(mels) | |
| f0 = self.decode(hiddens, thred=thred, use_viterbi=use_viterbi) | |
| f0s = [] | |
| for i in range(f0.shape[0]): | |
| f = f0[i, :sizes[i]] if sizes is not None else f0[i, :] | |
| f0s.append(f) | |
| return f0s | |
| def get_pitch_batch(self, waveforms, sample_rate, hop_size, lengths, interp_uv=False, fmin=50, fmax=1000): | |
| # hop_size, sample_rate: tgt params | |
| f0s = self.infer_from_audio_batch(waveforms, sample_rate=sample_rate) | |
| f0s_res, uvs_res = [], [] | |
| for idx, f0 in enumerate(f0s): | |
| f0 = self.postprocess(f0, fmin, fmax, min_gap=6) | |
| uv = f0 == 0 | |
| length = lengths[idx] | |
| time_step = hop_size / sample_rate | |
| f0_res = resample_align_curve(f0, 0.01, time_step, length) | |
| uv_res = resample_align_curve(uv.astype(np.float32), 0.01, time_step, length) > 0.5 | |
| if not interp_uv: | |
| f0_res[uv_res] = 0 | |
| f0s_res.append(f0_res) | |
| uvs_res.append(uv_res) | |
| return f0s_res, uvs_res | |
| def release_cuda(self): | |
| self.model = self.model.cpu() | |
| self.mel_extractor = self.mel_extractor.cpu() | |