Aloento/9Nine-PITS
1
1# modified from https://github.com/jaywalnut310/vits2import os3import random4 5import torch6import torch.utils.data7 8import commons9from analysis import Pitch10from mel_processing import spectrogram_torch11from text import cleaned_text_to_sequence12from utils import load_wav_to_torch, load_filepaths_and_text13 14""" Modified from Multi speaker version of VITS"""15 16 17class TextAudioSpeakerLoader(torch.utils.data.Dataset):18 """19 1) loads audio, speaker_id, text pairs20 2) normalizes text and converts them to sequences of integers21 3) computes spectrograms from audio files.22 """23 24 def __init__(self, audiopaths_sid_text, hparams, pt_run=False):25 self.audiopaths_sid_text = load_filepaths_and_text(audiopaths_sid_text)26 self.sampling_rate = hparams.sampling_rate27 self.filter_length = hparams.filter_length28 self.hop_length = hparams.hop_length29 self.win_length = hparams.win_length30 31 self.add_blank = hparams.add_blank32 self.min_text_len = 133 self.max_text_len = 19034 35 self.speaker_dict = {36 speaker: idx37 for idx, speaker in enumerate(hparams.speakers)38 }39 self.data_path = hparams.data_path40 41 self.pitch = Pitch(sr=hparams.sampling_rate,42 W=hparams.tau_max,43 tau_max=hparams.tau_max,44 midi_start=hparams.midi_start,45 midi_end=hparams.midi_end,46 octave_range=hparams.octave_range)47 48 random.seed(1234)49 random.shuffle(self.audiopaths_sid_text)50 self._filter()51 if pt_run:52 for _audiopaths_sid_text in self.audiopaths_sid_text:53 _ = self.get_audio_text_speaker_pair(_audiopaths_sid_text,54 True)55 56 def _filter(self):57 """58 Filter text & store spec lengths59 """60 # Store spectrogram lengths for Bucketing61 # wav_length ~= file_size / (wav_channels * Bytes per dim) = file_size / (1 * 2)62 # spec_length = wav_length // hop_length63 64 audiopaths_sid_text_new = []65 lengths = []66 for audiopath, spk, text, lang in self.audiopaths_sid_text:67 if self.min_text_len <= len(text) and len(68 text) <= self.max_text_len:69 audiopath = os.path.join(self.data_path, audiopath)70 if not os.path.exists(audiopath):71 print(audiopath, "not exist!")72 continue73 try:74 audio, sampling_rate = load_wav_to_torch(audiopath)75 except:76 print(audiopath, "load error!")77 continue78 audiopaths_sid_text_new.append([audiopath, spk, text, lang])79 lengths.append(80 os.path.getsize(audiopath) // (2 * self.hop_length))81 self.audiopaths_sid_text = audiopaths_sid_text_new82 self.lengths = lengths83 84 def get_audio_text_speaker_pair(self, audiopath_sid_text, pt_run=False):85 # separate filename, speaker_id and text86 audiopath, spk, text, lang = audiopath_sid_text87 text, lang = self.get_text(text, lang)88 spec, ying, wav = self.get_audio(audiopath, pt_run)89 sid = self.get_sid(self.speaker_dict[spk])90 return (text, spec, ying, wav, sid, lang)91 92 def get_audio(self, filename, pt_run=False):93 audio, sampling_rate = load_wav_to_torch(filename)94 if sampling_rate != self.sampling_rate:95 raise ValueError("{} {} SR doesn't match target {} SR".format(96 sampling_rate, self.sampling_rate))97 audio_norm = audio.unsqueeze(0)98 spec_filename = filename.replace(".wav", ".spec.pt")99 ying_filename = filename.replace(".wav", ".ying.pt")100 if os.path.exists(spec_filename) and not pt_run:101 spec = torch.load(spec_filename, map_location='cpu')102 else:103 spec = spectrogram_torch(audio_norm,104 self.filter_length,105 self.sampling_rate,106 self.hop_length,107 self.win_length,108 center=False)109 spec = torch.squeeze(spec, 0)110 torch.save(spec, spec_filename)111 if os.path.exists(ying_filename) and not pt_run:112 ying = torch.load(ying_filename, map_location='cpu')113 else:114 wav = torch.nn.functional.pad(115 audio_norm.unsqueeze(0),116 (self.filter_length - self.hop_length,117 self.filter_length - self.hop_length +118 (-audio_norm.shape[1]) % self.hop_length + self.hop_length * (audio_norm.shape[1] % self.hop_length == 0)),119 mode='constant').squeeze(0)120 ying = self.pitch.yingram(wav)[0]121 torch.save(ying, ying_filename)122 return spec, ying, audio_norm123 124 def get_text(self, text, lang):125 text_norm = cleaned_text_to_sequence(text)126 lang = [int(i) for i in lang.split(" ")]127 if self.add_blank:128 text_norm, lang = commons.intersperse_with_language_id(text_norm, lang, 0)129 text_norm = torch.LongTensor(text_norm)130 lang = torch.LongTensor(lang)131 return text_norm, lang132 133 def get_sid(self, sid):134 sid = torch.LongTensor([int(sid)])135 return sid136 137 def __getitem__(self, index):138 return self.get_audio_text_speaker_pair(139 self.audiopaths_sid_text[index])140 141 def __len__(self):142 return len(self.audiopaths_sid_text)143 144 145class TextAudioSpeakerCollate():146 """ Zero-pads model inputs and targets"""147 148 def __init__(self, return_ids=False):149 self.return_ids = return_ids150 151 def __call__(self, batch):152 """Collate's training batch from normalized text, audio and speaker identities153 PARAMS154 ------155 batch: [text_normalized, spec_normalized, wav_normalized, sid]156 """157 # Right zero-pad all one-hot text sequences to max input length158 _, ids_sorted_decreasing = torch.sort(torch.LongTensor(159 [x[1].size(1) for x in batch]),160 dim=0,161 descending=True)162 163 max_text_len = max([len(x[0]) for x in batch])164 max_spec_len = max([x[1].size(1) for x in batch])165 max_ying_len = max([x[2].size(1) for x in batch])166 max_wav_len = max([x[3].size(1) for x in batch])167 168 text_lengths = torch.LongTensor(len(batch))169 spec_lengths = torch.LongTensor(len(batch))170 ying_lengths = torch.LongTensor(len(batch))171 wav_lengths = torch.LongTensor(len(batch))172 sid = torch.LongTensor(len(batch))173 174 text_padded = torch.LongTensor(len(batch), max_text_len)175 tone_padded = torch.LongTensor(len(batch), max_text_len)176 spec_padded = torch.FloatTensor(len(batch), batch[0][1].size(0),177 max_spec_len)178 ying_padded = torch.FloatTensor(len(batch), batch[0][2].size(0),179 max_ying_len)180 wav_padded = torch.FloatTensor(len(batch), 1, max_wav_len)181 text_padded.zero_()182 tone_padded.zero_()183 spec_padded.zero_()184 ying_padded.zero_()185 wav_padded.zero_()186 for i in range(len(ids_sorted_decreasing)):187 row = batch[ids_sorted_decreasing[i]]188 189 text = row[0]190 text_padded[i, :text.size(0)] = text191 text_lengths[i] = text.size(0)192 193 spec = row[1]194 spec_padded[i, :, :spec.size(1)] = spec195 spec_lengths[i] = spec.size(1)196 197 ying = row[2]198 ying_padded[i, :, :ying.size(1)] = ying199 ying_lengths[i] = ying.size(1)200 201 wav = row[3]202 wav_padded[i, :, :wav.size(1)] = wav203 wav_lengths[i] = wav.size(1)204 205 tone = row[5]206 tone_padded[i, :text.size(0)] = tone207 208 sid[i] = row[4]209 210 if self.return_ids:211 return text_padded, text_lengths, spec_padded, spec_lengths, wav_padded, wav_lengths, sid, ids_sorted_decreasing212 return text_padded, text_lengths, spec_padded, spec_lengths, ying_padded, ying_lengths, wav_padded, wav_lengths, sid, tone_padded213 214 215class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler216 ):217 """218 Maintain similar input lengths in a batch.219 Length groups are specified by boundaries.220 Ex) boundaries = [b1, b2, b3] -> any batch is included either {x | b1 < length(x) <=b2} or {x | b2 < length(x) <= b3}.221 222 It removes samples which are not included in the boundaries.223 Ex) boundaries = [b1, b2, b3] -> any x s.t. length(x) <= b1 or length(x) > b3 are discarded.224 """225 226 def __init__(self,227 dataset,228 batch_size,229 boundaries,230 num_replicas=None,231 rank=None,232 shuffle=True):233 super().__init__(dataset,234 num_replicas=num_replicas,235 rank=rank,236 shuffle=shuffle)237 self.lengths = dataset.lengths238 self.batch_size = batch_size239 self.boundaries = boundaries240 241 self.buckets, self.num_samples_per_bucket = self._create_buckets()242 self.total_size = sum(self.num_samples_per_bucket)243 self.num_samples = self.total_size // self.num_replicas244 245 def _create_buckets(self):246 buckets = [[] for _ in range(len(self.boundaries) - 1)]247 for i in range(len(self.lengths)):248 length = self.lengths[i]249 idx_bucket = self._bisect(length)250 if idx_bucket != -1:251 buckets[idx_bucket].append(i)252 253 for i in range(len(buckets) - 1, -1, -1):254 if len(buckets[i]) == 0:255 buckets.pop(i)256 self.boundaries.pop(i + 1)257 258 num_samples_per_bucket = []259 for i in range(len(buckets)):260 len_bucket = len(buckets[i])261 total_batch_size = self.num_replicas * self.batch_size262 rem = (total_batch_size -263 (len_bucket % total_batch_size)) % total_batch_size264 num_samples_per_bucket.append(len_bucket + rem)265 return buckets, num_samples_per_bucket266 267 def __iter__(self):268 # deterministically shuffle based on epoch269 g = torch.Generator()270 g.manual_seed(self.epoch)271 272 indices = []273 if self.shuffle:274 for bucket in self.buckets:275 indices.append(276 torch.randperm(len(bucket), generator=g).tolist())277 else:278 for bucket in self.buckets:279 indices.append(list(range(len(bucket))))280 281 batches = []282 for i in range(len(self.buckets)):283 bucket = self.buckets[i]284 len_bucket = len(bucket)285 ids_bucket = indices[i]286 num_samples_bucket = self.num_samples_per_bucket[i]287 288 # add extra samples to make it evenly divisible289 rem = num_samples_bucket - len_bucket290 ids_bucket = ids_bucket + ids_bucket * \291 (rem // len_bucket) + ids_bucket[:(rem % len_bucket)]292 293 # subsample294 ids_bucket = ids_bucket[self.rank::self.num_replicas]295 296 # batching297 for j in range(len(ids_bucket) // self.batch_size):298 batch = [299 bucket[idx]300 for idx in ids_bucket[j * self.batch_size:(j + 1) *301 self.batch_size]302 ]303 batches.append(batch)304 305 if self.shuffle:306 batch_ids = torch.randperm(len(batches), generator=g).tolist()307 batches = [batches[i] for i in batch_ids]308 self.batches = batches309 310 assert len(self.batches) * self.batch_size == self.num_samples311 return iter(self.batches)312 313 def _bisect(self, x, lo=0, hi=None):314 if hi is None:315 hi = len(self.boundaries) - 1316 317 if hi > lo:318 mid = (hi + lo) // 2319 if self.boundaries[mid] < x and x <= self.boundaries[mid + 1]:320 return mid321 elif x <= self.boundaries[mid]:322 return self._bisect(x, lo, mid)323 else:324 return self._bisect(x, mid + 1, hi)325 else:326 return -1327 328 def __len__(self):329 return self.num_samples // self.batch_size330 331 332def create_spec(audiopaths_sid_text, hparams):333 audiopaths_sid_text = load_filepaths_and_text(audiopaths_sid_text)334 for audiopath, _, _, _ in audiopaths_sid_text:335 audiopath = os.path.join(hparams.data_path, audiopath)336 if not os.path.exists(audiopath):337 print(audiopath, "not exist!")338 continue339 try:340 audio, sampling_rate = load_wav_to_torch(audiopath)341 except:342 print(audiopath, "load error!")343 continue344 if sampling_rate != hparams.sampling_rate:345 raise ValueError("{} {} SR doesn't match target {} SR".format(346 sampling_rate, hparams.sampling_rate))347 audio_norm = audio.unsqueeze(0)348 specpath = audiopath.replace(".wav", ".spec.pt")349 350 if not os.path.exists(specpath):351 spec = spectrogram_torch(audio_norm,352 hparams.filter_length,353 hparams.sampling_rate,354 hparams.hop_length,355 hparams.win_length,356 center=False)357 spec = torch.squeeze(spec, 0)358 torch.save(spec, specpath)359 