emiyogo/clesss
0
1import time2import os3import random4import numpy as np5import torch6import torch.utils.data7 8import modules.commons as commons9import utils10from modules.mel_processing import spectrogram_torch, spec_to_mel_torch11from utils import load_wav_to_torch, load_filepaths_and_text12 13# import h5py14 15 16"""Multi speaker version"""17 18 19class TextAudioSpeakerLoader(torch.utils.data.Dataset):20 """21 1) loads audio, speaker_id, text pairs22 2) normalizes text and converts them to sequences of integers23 3) computes spectrograms from audio files.24 """25 26 def __init__(self, audiopaths, hparams):27 self.audiopaths = load_filepaths_and_text(audiopaths)28 self.max_wav_value = hparams.data.max_wav_value29 self.sampling_rate = hparams.data.sampling_rate30 self.filter_length = hparams.data.filter_length31 self.hop_length = hparams.data.hop_length32 self.win_length = hparams.data.win_length33 self.sampling_rate = hparams.data.sampling_rate34 self.use_sr = hparams.train.use_sr35 self.spec_len = hparams.train.max_speclen36 self.spk_map = hparams.spk37 38 random.seed(1234)39 random.shuffle(self.audiopaths)40 41 def get_audio(self, filename):42 filename = filename.replace("\\", "/")43 audio, sampling_rate = load_wav_to_torch(filename)44 if sampling_rate != self.sampling_rate:45 raise ValueError("{} SR doesn't match target {} SR".format(46 sampling_rate, self.sampling_rate))47 audio_norm = audio / self.max_wav_value48 audio_norm = audio_norm.unsqueeze(0)49 spec_filename = filename.replace(".wav", ".spec.pt")50 if os.path.exists(spec_filename):51 spec = torch.load(spec_filename)52 else:53 spec = spectrogram_torch(audio_norm, self.filter_length,54 self.sampling_rate, self.hop_length, self.win_length,55 center=False)56 spec = torch.squeeze(spec, 0)57 torch.save(spec, spec_filename)58 59 spk = filename.split("/")[-2]60 spk = torch.LongTensor([self.spk_map[spk]])61 62 f0 = np.load(filename + ".f0.npy")63 f0, uv = utils.interpolate_f0(f0)64 f0 = torch.FloatTensor(f0)65 uv = torch.FloatTensor(uv)66 67 c = torch.load(filename+ ".soft.pt")68 c = utils.repeat_expand_2d(c.squeeze(0), f0.shape[0])69 70 71 lmin = min(c.size(-1), spec.size(-1))72 assert abs(c.size(-1) - spec.size(-1)) < 3, (c.size(-1), spec.size(-1), f0.shape, filename)73 assert abs(audio_norm.shape[1]-lmin * self.hop_length) < 3 * self.hop_length74 spec, c, f0, uv = spec[:, :lmin], c[:, :lmin], f0[:lmin], uv[:lmin]75 audio_norm = audio_norm[:, :lmin * self.hop_length]76 # if spec.shape[1] < 30:77 # print("skip too short audio:", filename)78 # return None79 if spec.shape[1] > 800:80 start = random.randint(0, spec.shape[1]-800)81 end = start + 79082 spec, c, f0, uv = spec[:, start:end], c[:, start:end], f0[start:end], uv[start:end]83 audio_norm = audio_norm[:, start * self.hop_length : end * self.hop_length]84 85 return c, f0, spec, audio_norm, spk, uv86 87 def __getitem__(self, index):88 return self.get_audio(self.audiopaths[index][0])89 90 def __len__(self):91 return len(self.audiopaths)92 93 94class TextAudioCollate:95 96 def __call__(self, batch):97 batch = [b for b in batch if b is not None]98 99 input_lengths, ids_sorted_decreasing = torch.sort(100 torch.LongTensor([x[0].shape[1] for x in batch]),101 dim=0, descending=True)102 103 max_c_len = max([x[0].size(1) for x in batch])104 max_wav_len = max([x[3].size(1) for x in batch])105 106 lengths = torch.LongTensor(len(batch))107 108 c_padded = torch.FloatTensor(len(batch), batch[0][0].shape[0], max_c_len)109 f0_padded = torch.FloatTensor(len(batch), max_c_len)110 spec_padded = torch.FloatTensor(len(batch), batch[0][2].shape[0], max_c_len)111 wav_padded = torch.FloatTensor(len(batch), 1, max_wav_len)112 spkids = torch.LongTensor(len(batch), 1)113 uv_padded = torch.FloatTensor(len(batch), max_c_len)114 115 c_padded.zero_()116 spec_padded.zero_()117 f0_padded.zero_()118 wav_padded.zero_()119 uv_padded.zero_()120 121 for i in range(len(ids_sorted_decreasing)):122 row = batch[ids_sorted_decreasing[i]]123 124 c = row[0]125 c_padded[i, :, :c.size(1)] = c126 lengths[i] = c.size(1)127 128 f0 = row[1]129 f0_padded[i, :f0.size(0)] = f0130 131 spec = row[2]132 spec_padded[i, :, :spec.size(1)] = spec133 134 wav = row[3]135 wav_padded[i, :, :wav.size(1)] = wav136 137 spkids[i, 0] = row[4]138 139 uv = row[5]140 uv_padded[i, :uv.size(0)] = uv141 142 return c_padded, f0_padded, spec_padded, wav_padded, spkids, lengths, uv_padded143 