Initial commit
This commit is contained in:
@@ -0,0 +1 @@
|
||||
"""Pitch extractor modules for ROSVOT."""
|
||||
@@ -0,0 +1,6 @@
|
||||
from .constants import *
|
||||
from .model import E2E0
|
||||
from .utils import to_local_average_f0, to_viterbi_f0
|
||||
from .inference import RMVPE
|
||||
from .spec import MelSpectrogram
|
||||
from .extractor import extract
|
||||
@@ -0,0 +1,9 @@
|
||||
SAMPLE_RATE = 16000
|
||||
|
||||
N_CLASS = 360
|
||||
|
||||
N_MELS = 128
|
||||
MEL_FMIN = 30
|
||||
MEL_FMAX = 8000
|
||||
WINDOW_LENGTH = 1024
|
||||
CONST = 1997.3794084376191
|
||||
@@ -0,0 +1,173 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from .constants import N_MELS
|
||||
|
||||
|
||||
class ConvBlockRes(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, momentum=0.01):
|
||||
super(ConvBlockRes, self).__init__()
|
||||
self.conv = nn.Sequential(
|
||||
nn.Conv2d(in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=(3, 3),
|
||||
stride=(1, 1),
|
||||
padding=(1, 1),
|
||||
bias=False),
|
||||
nn.BatchNorm2d(out_channels, momentum=momentum),
|
||||
nn.ReLU(),
|
||||
|
||||
nn.Conv2d(in_channels=out_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=(3, 3),
|
||||
stride=(1, 1),
|
||||
padding=(1, 1),
|
||||
bias=False),
|
||||
nn.BatchNorm2d(out_channels, momentum=momentum),
|
||||
nn.ReLU(),
|
||||
)
|
||||
if in_channels != out_channels:
|
||||
self.shortcut = nn.Conv2d(in_channels, out_channels, (1, 1))
|
||||
self.is_shortcut = True
|
||||
else:
|
||||
self.is_shortcut = False
|
||||
|
||||
def forward(self, x):
|
||||
if self.is_shortcut:
|
||||
return self.conv(x) + self.shortcut(x)
|
||||
else:
|
||||
return self.conv(x) + x
|
||||
|
||||
|
||||
class ResEncoderBlock(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, kernel_size, n_blocks=1, momentum=0.01):
|
||||
super(ResEncoderBlock, self).__init__()
|
||||
self.n_blocks = n_blocks
|
||||
self.conv = nn.ModuleList()
|
||||
self.conv.append(ConvBlockRes(in_channels, out_channels, momentum))
|
||||
for i in range(n_blocks - 1):
|
||||
self.conv.append(ConvBlockRes(out_channels, out_channels, momentum))
|
||||
self.kernel_size = kernel_size
|
||||
if self.kernel_size is not None:
|
||||
self.pool = nn.AvgPool2d(kernel_size=kernel_size)
|
||||
|
||||
def forward(self, x):
|
||||
for i in range(self.n_blocks):
|
||||
x = self.conv[i](x)
|
||||
if self.kernel_size is not None:
|
||||
return x, self.pool(x)
|
||||
else:
|
||||
return x
|
||||
|
||||
|
||||
class ResDecoderBlock(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, stride, n_blocks=1, momentum=0.01):
|
||||
super(ResDecoderBlock, self).__init__()
|
||||
out_padding = (0, 1) if stride == (1, 2) else (1, 1)
|
||||
self.n_blocks = n_blocks
|
||||
self.conv1 = nn.Sequential(
|
||||
nn.ConvTranspose2d(in_channels=in_channels,
|
||||
out_channels=out_channels,
|
||||
kernel_size=(3, 3),
|
||||
stride=stride,
|
||||
padding=(1, 1),
|
||||
output_padding=out_padding,
|
||||
bias=False),
|
||||
nn.BatchNorm2d(out_channels, momentum=momentum),
|
||||
nn.ReLU(),
|
||||
)
|
||||
self.conv2 = nn.ModuleList()
|
||||
self.conv2.append(ConvBlockRes(out_channels * 2, out_channels, momentum))
|
||||
for i in range(n_blocks-1):
|
||||
self.conv2.append(ConvBlockRes(out_channels, out_channels, momentum))
|
||||
|
||||
def forward(self, x, concat_tensor):
|
||||
x = self.conv1(x)
|
||||
x = torch.cat((x, concat_tensor), dim=1)
|
||||
for i in range(self.n_blocks):
|
||||
x = self.conv2[i](x)
|
||||
return x
|
||||
|
||||
|
||||
class Encoder(nn.Module):
|
||||
def __init__(self, in_channels, in_size, n_encoders, kernel_size, n_blocks, out_channels=16, momentum=0.01):
|
||||
super(Encoder, self).__init__()
|
||||
self.n_encoders = n_encoders
|
||||
self.bn = nn.BatchNorm2d(in_channels, momentum=momentum)
|
||||
self.layers = nn.ModuleList()
|
||||
self.latent_channels = []
|
||||
for i in range(self.n_encoders):
|
||||
self.layers.append(ResEncoderBlock(in_channels, out_channels, kernel_size, n_blocks, momentum=momentum))
|
||||
self.latent_channels.append([out_channels, in_size])
|
||||
in_channels = out_channels
|
||||
out_channels *= 2
|
||||
in_size //= 2
|
||||
self.out_size = in_size
|
||||
self.out_channel = out_channels
|
||||
|
||||
def forward(self, x):
|
||||
concat_tensors = []
|
||||
x = self.bn(x)
|
||||
for i in range(self.n_encoders):
|
||||
_, x = self.layers[i](x)
|
||||
concat_tensors.append(_)
|
||||
return x, concat_tensors
|
||||
|
||||
|
||||
class Intermediate(nn.Module):
|
||||
def __init__(self, in_channels, out_channels, n_inters, n_blocks, momentum=0.01):
|
||||
super(Intermediate, self).__init__()
|
||||
self.n_inters = n_inters
|
||||
self.layers = nn.ModuleList()
|
||||
self.layers.append(ResEncoderBlock(in_channels, out_channels, None, n_blocks, momentum))
|
||||
for i in range(self.n_inters-1):
|
||||
self.layers.append(ResEncoderBlock(out_channels, out_channels, None, n_blocks, momentum))
|
||||
|
||||
def forward(self, x):
|
||||
for i in range(self.n_inters):
|
||||
x = self.layers[i](x)
|
||||
return x
|
||||
|
||||
|
||||
class Decoder(nn.Module):
|
||||
def __init__(self, in_channels, n_decoders, stride, n_blocks, momentum=0.01):
|
||||
super(Decoder, self).__init__()
|
||||
self.layers = nn.ModuleList()
|
||||
self.n_decoders = n_decoders
|
||||
for i in range(self.n_decoders):
|
||||
out_channels = in_channels // 2
|
||||
self.layers.append(ResDecoderBlock(in_channels, out_channels, stride, n_blocks, momentum))
|
||||
in_channels = out_channels
|
||||
|
||||
def forward(self, x, concat_tensors):
|
||||
for i in range(self.n_decoders):
|
||||
x = self.layers[i](x, concat_tensors[-1-i])
|
||||
return x
|
||||
|
||||
|
||||
class TimbreFilter(nn.Module):
|
||||
def __init__(self, latent_rep_channels):
|
||||
super(TimbreFilter, self).__init__()
|
||||
self.layers = nn.ModuleList()
|
||||
for latent_rep in latent_rep_channels:
|
||||
self.layers.append(ConvBlockRes(latent_rep[0], latent_rep[0]))
|
||||
|
||||
def forward(self, x_tensors):
|
||||
out_tensors = []
|
||||
for i, layer in enumerate(self.layers):
|
||||
out_tensors.append(layer(x_tensors[i]))
|
||||
return out_tensors
|
||||
|
||||
|
||||
class DeepUnet0(nn.Module):
|
||||
def __init__(self, kernel_size, n_blocks, en_de_layers=5, inter_layers=4, in_channels=1, en_out_channels=16):
|
||||
super(DeepUnet0, self).__init__()
|
||||
self.encoder = Encoder(in_channels, N_MELS, en_de_layers, kernel_size, n_blocks, en_out_channels)
|
||||
self.intermediate = Intermediate(self.encoder.out_channel // 2, self.encoder.out_channel, inter_layers, n_blocks)
|
||||
self.tf = TimbreFilter(self.encoder.latent_channels)
|
||||
self.decoder = Decoder(self.encoder.out_channel, en_de_layers, kernel_size, n_blocks)
|
||||
|
||||
def forward(self, x):
|
||||
x, concat_tensors = self.encoder(x)
|
||||
x = self.intermediate(x)
|
||||
x = self.decoder(x, concat_tensors)
|
||||
return x
|
||||
@@ -0,0 +1,183 @@
|
||||
import math
|
||||
import os
|
||||
|
||||
from tqdm import tqdm
|
||||
import librosa
|
||||
import numpy as np
|
||||
import torch
|
||||
from torch.utils.data import Dataset, DataLoader, DistributedSampler
|
||||
import torch.multiprocessing as mp
|
||||
from torch.distributed import init_process_group
|
||||
import torch.distributed as dist
|
||||
|
||||
from .inference import RMVPE
|
||||
from ....utils.commons.dataset_utils import batch_by_size, build_dataloader
|
||||
# import utils
|
||||
from ....utils.audio import get_wav_num_frames
|
||||
|
||||
"""
|
||||
A convenient API for batch inference
|
||||
update: add ddp
|
||||
"""
|
||||
|
||||
class RMVPEInferDataset(Dataset):
|
||||
def __init__(self, wav_fns: list, id_and_sizes=None, sr=24000, hop_size=128, num_workers=0):
|
||||
if id_and_sizes is None:
|
||||
id_and_sizes = []
|
||||
if type(wav_fns[0]) == str: # wav_paths
|
||||
for idx, wav_path in enumerate(wav_fns):
|
||||
total_frames = get_wav_num_frames(wav_path, sr)
|
||||
id_and_sizes.append((idx, round(total_frames / hop_size)))
|
||||
else: # numpy arrays, mono wavs
|
||||
for idx, wav in enumerate(wav_fns):
|
||||
id_and_sizes.append((idx, round(wav.shape[-1] / hop_size)))
|
||||
self.wav_fns = wav_fns
|
||||
self.id_and_sizes = id_and_sizes
|
||||
self.sr = sr
|
||||
self.num_workers = num_workers
|
||||
|
||||
def __getitem__(self, idx):
|
||||
if type(self.wav_fns[idx]) == str:
|
||||
wav_fn = self.wav_fns[idx]
|
||||
wav, _ = librosa.core.load(wav_fn, sr=self.sr)
|
||||
else:
|
||||
wav = self.wav_fns[idx]
|
||||
return idx, wav
|
||||
|
||||
def collater(self, samples: list):
|
||||
return samples
|
||||
|
||||
def __len__(self):
|
||||
return len(self.wav_fns)
|
||||
|
||||
def ordered_indices(self):
|
||||
"""Return an ordered list of indices. Batches will be constructed based
|
||||
on this order."""
|
||||
return np.arange(len(self))
|
||||
|
||||
def num_tokens(self, index):
|
||||
return self.id_and_sizes[index][1]
|
||||
|
||||
@torch.no_grad()
|
||||
def extract(wav_fns: list, id_and_sizes=None, ckpt=None, sr=24000, hop_size=128, bsz=128, max_tokens=100000,
|
||||
fmax=900, fmin=50, ds_workers=0):
|
||||
all_gpu_ids = [int(x) for x in os.environ.get("CUDA_VISIBLE_DEVICES", "").split(",") if x != '']
|
||||
num_gpus = len(all_gpu_ids)
|
||||
dist_config = {
|
||||
"dist_backend": "nccl",
|
||||
"dist_url": "tcp://localhost:54189",
|
||||
"world_size": 1
|
||||
}
|
||||
# https://discuss.pytorch.org/t/how-to-fix-a-sigsegv-in-pytorch-when-using-distributed-training-e-g-ddp/113518/10#:~:text=Using%20start%20and%20join%20avoids
|
||||
# https://github.com/pytorch/pytorch/issues/40403#issuecomment-648515174
|
||||
# mp.set_start_method('spawn')
|
||||
if num_gpus > 1:
|
||||
result_queue = mp.Queue()
|
||||
for rank in range(num_gpus):
|
||||
mp.Process(target=extract_worker, args=(rank, wav_fns, id_and_sizes, ckpt, sr, hop_size, bsz, max_tokens, fmax,
|
||||
fmin, dist_config, num_gpus, ds_workers, result_queue,)).start()
|
||||
f0_res = [None] * len(wav_fns)
|
||||
for _ in range(num_gpus):
|
||||
f0_res_dict = result_queue.get()
|
||||
for idx in f0_res_dict:
|
||||
f0_res[idx] = f0_res_dict[idx]
|
||||
del f0_res_dict
|
||||
else:
|
||||
# f0_res = extract_one_process(wav_fns, id_and_sizes, ckpt, sr, hop_size, bsz, max_tokens, fmax, fmin)
|
||||
f0_res_dict = extract_worker(0, wav_fns, id_and_sizes, ckpt, sr, hop_size, bsz, max_tokens, fmax,
|
||||
fmin, dist_config, num_gpus, ds_workers, None)
|
||||
f0_res = [None] * len(wav_fns)
|
||||
for idx in f0_res_dict:
|
||||
f0_res[idx] = f0_res_dict[idx]
|
||||
return f0_res
|
||||
|
||||
@torch.no_grad()
|
||||
def extract_worker(rank, wav_fns: list, id_and_sizes=None, ckpt=None, sr=24000, hop_size=128, bsz=128, max_tokens=100000,
|
||||
fmax=900, fmin=50, dist_config=None, num_gpus=1, ds_workers=0, q=None):
|
||||
# print(f"rank: {rank}")
|
||||
if num_gpus > 1:
|
||||
init_process_group(backend=dist_config['dist_backend'], init_method=dist_config['dist_url'],
|
||||
world_size=dist_config['world_size'] * num_gpus, rank=rank)
|
||||
dataset = RMVPEInferDataset(wav_fns, id_and_sizes, sr, hop_size, num_workers=ds_workers)
|
||||
# ds_sampler = DistributedSampler(dataset, shuffle=False) if num_gpus > 1 else None
|
||||
# loader = DataLoader(dataset, sampler=ds_sampler, collate_fn=dataset.collator, batch_size=1, num_workers=40, drop_last=False)
|
||||
loader = build_dataloader(dataset, shuffle=False, max_tokens=max_tokens, max_sentences=bsz, use_ddp=num_gpus > 1)
|
||||
loader = tqdm(loader, desc=f'| Processing f0 in [n_ranks={num_gpus}; max_tokens={max_tokens}; max_sentences={bsz}]') if rank == 0 else loader
|
||||
|
||||
device = torch.device(f"cuda:{int(rank)}")
|
||||
model = RMVPE(ckpt, device=device)
|
||||
f0_res_dict = {}
|
||||
for batch in loader:
|
||||
if batch is None or len(batch) == 0:
|
||||
continue
|
||||
idxs = [item[0] for item in batch]
|
||||
wavs = [item[1] for item in batch]
|
||||
lengths = [(wav.shape[0] + hop_size - 1) // hop_size for wav in wavs]
|
||||
with torch.no_grad():
|
||||
f0s, uvs = model.get_pitch_batch(
|
||||
wavs, sample_rate=sr,
|
||||
hop_size=hop_size,
|
||||
lengths=lengths,
|
||||
fmax=fmax,
|
||||
fmin=fmin
|
||||
)
|
||||
for i, idx in enumerate(idxs):
|
||||
f0_res_dict[idx] = f0s[i]
|
||||
if q is not None:
|
||||
q.put(f0_res_dict)
|
||||
else:
|
||||
return f0_res_dict
|
||||
|
||||
# old version
|
||||
def extract_one_process(wav_fns: list, id_and_sizes=None, ckpt=None, sr=24000, hop_size=128, bsz=128, max_tokens=100000,
|
||||
fmax=900, fmin=50, device='cuda'):
|
||||
assert ckpt is not None
|
||||
rmvpe = RMVPE(ckpt, device=device)
|
||||
if id_and_sizes is None:
|
||||
id_and_sizes = []
|
||||
if type(wav_fns[0]) == str: # wav_paths
|
||||
for idx, wav_path in enumerate(wav_fns):
|
||||
total_frames = get_wav_num_frames(wav_path, sr)
|
||||
id_and_sizes.append((idx, round(total_frames / hop_size)))
|
||||
else: # numpy arrays, mono wavs
|
||||
for idx, wav in enumerate(wav_fns):
|
||||
id_and_sizes.append((idx, round(wav.shape[-1] / hop_size)))
|
||||
get_size = lambda x: x[1]
|
||||
bs = batch_by_size(id_and_sizes, get_size, max_tokens=max_tokens, max_sentences=bsz)
|
||||
for i in range(len(bs)):
|
||||
bs[i] = [bs[i][j][0] for j in range(len(bs[i]))]
|
||||
|
||||
f0_res = [None] * len(wav_fns)
|
||||
for batch in tqdm(bs, total=len(bs), desc=f'| Processing f0 in [max_tokens={max_tokens}; max_sentences={bsz}]'):
|
||||
wavs, mel_lengths, lengths = [], [], []
|
||||
for idx in batch:
|
||||
if type(wav_fns[idx]) == str:
|
||||
wav_fn = wav_fns[idx]
|
||||
wav, _ = librosa.core.load(wav_fn, sr=sr)
|
||||
else:
|
||||
wav = wav_fns[idx]
|
||||
wavs.append(wav)
|
||||
mel_lengths.append(math.ceil((wav.shape[0] + 1) / hop_size))
|
||||
lengths.append((wav.shape[0] + hop_size - 1) // hop_size)
|
||||
|
||||
with torch.no_grad():
|
||||
f0s, uvs = rmvpe.get_pitch_batch(
|
||||
wavs, sample_rate=sr,
|
||||
hop_size=hop_size,
|
||||
lengths=lengths,
|
||||
fmax=fmax,
|
||||
fmin=fmin
|
||||
)
|
||||
|
||||
for i, idx in enumerate(batch):
|
||||
f0_res[idx] = f0s[i]
|
||||
|
||||
if rmvpe is not None:
|
||||
rmvpe.release_cuda()
|
||||
torch.cuda.empty_cache()
|
||||
rmvpe = None
|
||||
|
||||
return f0_res
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
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
|
||||
|
||||
@torch.no_grad()
|
||||
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()
|
||||
@@ -0,0 +1,32 @@
|
||||
from torch import nn
|
||||
|
||||
from .constants import *
|
||||
from .deepunet import DeepUnet0
|
||||
from .seq import BiGRU
|
||||
|
||||
|
||||
class E2E0(nn.Module):
|
||||
def __init__(self, n_blocks, n_gru, kernel_size, en_de_layers=5, inter_layers=4, in_channels=1,
|
||||
en_out_channels=16):
|
||||
super(E2E0, self).__init__()
|
||||
self.unet = DeepUnet0(kernel_size, n_blocks, en_de_layers, inter_layers, in_channels, en_out_channels)
|
||||
self.cnn = nn.Conv2d(en_out_channels, 3, (3, 3), padding=(1, 1))
|
||||
if n_gru:
|
||||
self.fc = nn.Sequential(
|
||||
BiGRU(3 * N_MELS, 256, n_gru),
|
||||
nn.Linear(512, N_CLASS),
|
||||
nn.Dropout(0.25),
|
||||
nn.Sigmoid()
|
||||
)
|
||||
else:
|
||||
self.fc = nn.Sequential(
|
||||
nn.Linear(3 * N_MELS, N_CLASS),
|
||||
nn.Dropout(0.25),
|
||||
nn.Sigmoid()
|
||||
)
|
||||
|
||||
def forward(self, mel):
|
||||
mel = mel.transpose(-1, -2).unsqueeze(1)
|
||||
x = self.cnn(self.unet(mel)).transpose(1, 2).flatten(-2)
|
||||
x = self.fc(x)
|
||||
return x
|
||||
@@ -0,0 +1,10 @@
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class BiGRU(nn.Module):
|
||||
def __init__(self, input_features, hidden_features, num_layers):
|
||||
super(BiGRU, self).__init__()
|
||||
self.gru = nn.GRU(input_features, hidden_features, num_layers=num_layers, batch_first=True, bidirectional=True)
|
||||
|
||||
def forward(self, x):
|
||||
return self.gru(x)[0]
|
||||
@@ -0,0 +1,72 @@
|
||||
import torch
|
||||
import numpy as np
|
||||
import torch.nn.functional as F
|
||||
from librosa.filters import mel
|
||||
|
||||
|
||||
class MelSpectrogram(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
n_mel_channels,
|
||||
sampling_rate,
|
||||
win_length,
|
||||
hop_length,
|
||||
n_fft=None,
|
||||
mel_fmin=0,
|
||||
mel_fmax=None,
|
||||
clamp=1e-5
|
||||
):
|
||||
super().__init__()
|
||||
n_fft = win_length if n_fft is None else n_fft
|
||||
self.hann_window = {}
|
||||
mel_basis = mel(
|
||||
sr=sampling_rate,
|
||||
n_fft=n_fft,
|
||||
n_mels=n_mel_channels,
|
||||
fmin=mel_fmin,
|
||||
fmax=mel_fmax,
|
||||
htk=True)
|
||||
mel_basis = torch.from_numpy(mel_basis).float()
|
||||
self.register_buffer("mel_basis", mel_basis)
|
||||
self.n_fft = win_length if n_fft is None else n_fft
|
||||
self.hop_length = hop_length
|
||||
self.win_length = win_length
|
||||
self.sampling_rate = sampling_rate
|
||||
self.n_mel_channels = n_mel_channels
|
||||
self.clamp = clamp
|
||||
|
||||
def forward(self, audio, keyshift=0, speed=1, center=True):
|
||||
factor = 2 ** (keyshift / 12)
|
||||
n_fft_new = int(np.round(self.n_fft * factor))
|
||||
win_length_new = int(np.round(self.win_length * factor))
|
||||
hop_length_new = int(np.round(self.hop_length * speed))
|
||||
|
||||
keyshift_key = str(keyshift) + '_' + str(audio.device)
|
||||
if keyshift_key not in self.hann_window:
|
||||
self.hann_window[keyshift_key] = torch.hann_window(win_length_new).to(audio.device)
|
||||
if center:
|
||||
pad_left = win_length_new // 2
|
||||
pad_right = (win_length_new + 1) // 2
|
||||
audio = F.pad(audio, (pad_left, pad_right))
|
||||
|
||||
fft = torch.stft(
|
||||
audio,
|
||||
n_fft=n_fft_new,
|
||||
hop_length=hop_length_new,
|
||||
win_length=win_length_new,
|
||||
window=self.hann_window[keyshift_key],
|
||||
center=False,
|
||||
return_complex=True
|
||||
)
|
||||
magnitude = fft.abs()
|
||||
|
||||
if keyshift != 0:
|
||||
size = self.n_fft // 2 + 1
|
||||
resize = magnitude.size(1)
|
||||
if resize < size:
|
||||
magnitude = F.pad(magnitude, (0, 0, 0, size - resize))
|
||||
magnitude = magnitude[:, :size, :] * self.win_length / win_length_new
|
||||
|
||||
mel_output = torch.matmul(self.mel_basis, magnitude)
|
||||
log_mel_spec = torch.log(torch.clamp(mel_output, min=self.clamp))
|
||||
return log_mel_spec
|
||||
@@ -0,0 +1,43 @@
|
||||
import librosa
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from .constants import *
|
||||
|
||||
|
||||
def to_local_average_f0(hidden, center=None, thred=0.03):
|
||||
idx = torch.arange(N_CLASS, device=hidden.device)[None, None, :] # [B=1, T=1, N]
|
||||
idx_cents = idx * 20 + CONST # [B=1, N]
|
||||
if center is None:
|
||||
center = torch.argmax(hidden, dim=2, keepdim=True) # [B, T, 1]
|
||||
start = torch.clip(center - 4, min=0) # [B, T, 1]
|
||||
end = torch.clip(center + 5, max=N_CLASS) # [B, T, 1]
|
||||
idx_mask = (idx >= start) & (idx < end) # [B, T, N]
|
||||
weights = hidden * idx_mask # [B, T, N]
|
||||
product_sum = torch.sum(weights * idx_cents, dim=2) # [B, T]
|
||||
weight_sum = torch.sum(weights, dim=2) # [B, T]
|
||||
cents = product_sum / (weight_sum + (weight_sum == 0)) # avoid dividing by zero, [B, T]
|
||||
f0 = 10 * 2 ** (cents / 1200)
|
||||
uv = hidden.max(dim=2)[0] < thred # [B, T]
|
||||
f0 = f0 * ~uv
|
||||
return f0.cpu().numpy()
|
||||
|
||||
|
||||
def to_viterbi_f0(hidden, thred=0.03):
|
||||
# Create viterbi transition matrix
|
||||
if not hasattr(to_viterbi_f0, 'transition'):
|
||||
xx, yy = np.meshgrid(range(N_CLASS), range(N_CLASS))
|
||||
transition = np.maximum(30 - abs(xx - yy), 0)
|
||||
transition = transition / transition.sum(axis=1, keepdims=True)
|
||||
to_viterbi_f0.transition = transition
|
||||
|
||||
# Convert to probability
|
||||
prob = hidden.squeeze(0).cpu().numpy()
|
||||
prob = prob.T
|
||||
prob = prob / prob.sum(axis=0)
|
||||
|
||||
# Perform viterbi decoding
|
||||
path = librosa.sequence.viterbi(prob, to_viterbi_f0.transition).astype(np.int64)
|
||||
center = torch.from_numpy(path).unsqueeze(0).unsqueeze(-1).to(hidden.device)
|
||||
|
||||
return to_local_average_f0(hidden, center=center, thred=thred)
|
||||
Reference in New Issue
Block a user