Kafke/Code-Realize-TTS
0
1import time2import os3import random4import numpy as np5import torch6import torch.utils.data7 8import commons 9from mel_processing import spectrogram_torch10from utils import load_wav_to_torch, load_filepaths_and_text11from text import text_to_sequence, cleaned_text_to_sequence12 13 14class TextAudioLoader(torch.utils.data.Dataset):15 """16 1) loads audio, text pairs17 2) normalizes text and converts them to sequences of integers18 3) computes spectrograms from audio files.19 """20 def __init__(self, audiopaths_and_text, hparams):21 self.audiopaths_and_text = load_filepaths_and_text(audiopaths_and_text)22 self.text_cleaners = hparams.text_cleaners23 self.max_wav_value = hparams.max_wav_value24 self.sampling_rate = hparams.sampling_rate25 self.filter_length = hparams.filter_length 26 self.hop_length = hparams.hop_length 27 self.win_length = hparams.win_length28 self.sampling_rate = hparams.sampling_rate 29 30 self.cleaned_text = getattr(hparams, "cleaned_text", False)31 32 self.add_blank = hparams.add_blank33 self.min_text_len = getattr(hparams, "min_text_len", 1)34 self.max_text_len = getattr(hparams, "max_text_len", 190)35 36 random.seed(1234)37 random.shuffle(self.audiopaths_and_text)38 self._filter()39 40 41 def _filter(self):42 """43 Filter text & store spec lengths44 """45 # Store spectrogram lengths for Bucketing46 # wav_length ~= file_size / (wav_channels * Bytes per dim) = file_size / (1 * 2)47 # spec_length = wav_length // hop_length48 49 audiopaths_and_text_new = []50 lengths = []51 for audiopath, text in self.audiopaths_and_text:52 if self.min_text_len <= len(text) and len(text) <= self.max_text_len:53 audiopaths_and_text_new.append([audiopath, text])54 lengths.append(os.path.getsize(audiopath) // (2 * self.hop_length))55 self.audiopaths_and_text = audiopaths_and_text_new56 self.lengths = lengths57 58 def get_audio_text_pair(self, audiopath_and_text):59 # separate filename and text60 audiopath, text = audiopath_and_text[0], audiopath_and_text[1]61 text = self.get_text(text)62 spec, wav = self.get_audio(audiopath)63 return (text, spec, wav)64 65 def get_audio(self, filename):66 audio, sampling_rate = load_wav_to_torch(filename)67 if sampling_rate != self.sampling_rate:68 raise ValueError("{} {} SR doesn't match target {} SR".format(69 sampling_rate, self.sampling_rate))70 audio_norm = audio / self.max_wav_value71 audio_norm = audio_norm.unsqueeze(0)72 spec_filename = filename.replace(".wav", ".spec.pt")73 if os.path.exists(spec_filename):74 spec = torch.load(spec_filename)75 else:76 spec = spectrogram_torch(audio_norm, self.filter_length,77 self.sampling_rate, self.hop_length, self.win_length,78 center=False)79 spec = torch.squeeze(spec, 0)80 torch.save(spec, spec_filename)81 return spec, audio_norm82 83 def get_text(self, text):84 if self.cleaned_text:85 text_norm = cleaned_text_to_sequence(text)86 else:87 text_norm = text_to_sequence(text, self.text_cleaners)88 if self.add_blank:89 text_norm = commons.intersperse(text_norm, 0)90 text_norm = torch.LongTensor(text_norm)91 return text_norm92 93 def __getitem__(self, index):94 return self.get_audio_text_pair(self.audiopaths_and_text[index])95 96 def __len__(self):97 return len(self.audiopaths_and_text)98 99 100class TextAudioCollate():101 """ Zero-pads model inputs and targets102 """103 def __init__(self, return_ids=False):104 self.return_ids = return_ids105 106 def __call__(self, batch):107 """Collate's training batch from normalized text and aduio108 PARAMS109 ------110 batch: [text_normalized, spec_normalized, wav_normalized]111 """112 # Right zero-pad all one-hot text sequences to max input length113 _, ids_sorted_decreasing = torch.sort(114 torch.LongTensor([x[1].size(1) for x in batch]),115 dim=0, descending=True)116 117 max_text_len = max([len(x[0]) for x in batch])118 max_spec_len = max([x[1].size(1) for x in batch])119 max_wav_len = max([x[2].size(1) for x in batch])120 121 text_lengths = torch.LongTensor(len(batch))122 spec_lengths = torch.LongTensor(len(batch))123 wav_lengths = torch.LongTensor(len(batch))124 125 text_padded = torch.LongTensor(len(batch), max_text_len)126 spec_padded = torch.FloatTensor(len(batch), batch[0][1].size(0), max_spec_len)127 wav_padded = torch.FloatTensor(len(batch), 1, max_wav_len)128 text_padded.zero_()129 spec_padded.zero_()130 wav_padded.zero_()131 for i in range(len(ids_sorted_decreasing)):132 row = batch[ids_sorted_decreasing[i]]133 134 text = row[0]135 text_padded[i, :text.size(0)] = text136 text_lengths[i] = text.size(0)137 138 spec = row[1]139 spec_padded[i, :, :spec.size(1)] = spec140 spec_lengths[i] = spec.size(1)141 142 wav = row[2]143 wav_padded[i, :, :wav.size(1)] = wav144 wav_lengths[i] = wav.size(1)145 146 if self.return_ids:147 return text_padded, text_lengths, spec_padded, spec_lengths, wav_padded, wav_lengths, ids_sorted_decreasing148 return text_padded, text_lengths, spec_padded, spec_lengths, wav_padded, wav_lengths149 150 151"""Multi speaker version"""152class TextAudioSpeakerLoader(torch.utils.data.Dataset):153 """154 1) loads audio, speaker_id, text pairs155 2) normalizes text and converts them to sequences of integers156 3) computes spectrograms from audio files.157 """158 def __init__(self, audiopaths_sid_text, hparams):159 self.audiopaths_sid_text = load_filepaths_and_text(audiopaths_sid_text)160 self.text_cleaners = hparams.text_cleaners161 self.max_wav_value = hparams.max_wav_value162 self.sampling_rate = hparams.sampling_rate163 self.filter_length = hparams.filter_length164 self.hop_length = hparams.hop_length165 self.win_length = hparams.win_length166 self.sampling_rate = hparams.sampling_rate167 168 self.cleaned_text = getattr(hparams, "cleaned_text", False)169 170 self.add_blank = hparams.add_blank171 self.min_text_len = getattr(hparams, "min_text_len", 1)172 self.max_text_len = getattr(hparams, "max_text_len", 190)173 174 random.seed(1234)175 random.shuffle(self.audiopaths_sid_text)176 self._filter()177 178 def _filter(self):179 """180 Filter text & store spec lengths181 """182 # Store spectrogram lengths for Bucketing183 # wav_length ~= file_size / (wav_channels * Bytes per dim) = file_size / (1 * 2)184 # spec_length = wav_length // hop_length185 186 audiopaths_sid_text_new = []187 lengths = []188 for audiopath, sid, text in self.audiopaths_sid_text:189 audiopath = "E:/uma_voice/" + audiopath190 if self.min_text_len <= len(text) and len(text) <= self.max_text_len:191 audiopaths_sid_text_new.append([audiopath, sid, text])192 lengths.append(os.path.getsize(audiopath) // (2 * self.hop_length))193 self.audiopaths_sid_text = audiopaths_sid_text_new194 self.lengths = lengths195 196 def get_audio_text_speaker_pair(self, audiopath_sid_text):197 # separate filename, speaker_id and text198 audiopath, sid, text = audiopath_sid_text[0], audiopath_sid_text[1], audiopath_sid_text[2]199 text = self.get_text(text)200 spec, wav = self.get_audio(audiopath)201 sid = self.get_sid(sid)202 return (text, spec, wav, sid)203 204 def get_audio(self, filename):205 audio, sampling_rate = load_wav_to_torch(filename)206 if sampling_rate != self.sampling_rate:207 raise ValueError("{} {} SR doesn't match target {} SR".format(208 sampling_rate, self.sampling_rate))209 audio_norm = audio / self.max_wav_value210 audio_norm = audio_norm.unsqueeze(0)211 spec_filename = filename.replace(".wav", ".spec.pt")212 if os.path.exists(spec_filename):213 spec = torch.load(spec_filename)214 else:215 spec = spectrogram_torch(audio_norm, self.filter_length,216 self.sampling_rate, self.hop_length, self.win_length,217 center=False)218 spec = torch.squeeze(spec, 0)219 torch.save(spec, spec_filename)220 return spec, audio_norm221 222 def get_text(self, text):223 if self.cleaned_text:224 text_norm = cleaned_text_to_sequence(text)225 else:226 text_norm = text_to_sequence(text, self.text_cleaners)227 if self.add_blank:228 text_norm = commons.intersperse(text_norm, 0)229 text_norm = torch.LongTensor(text_norm)230 return text_norm231 232 def get_sid(self, sid):233 sid = torch.LongTensor([int(sid)])234 return sid235 236 def __getitem__(self, index):237 return self.get_audio_text_speaker_pair(self.audiopaths_sid_text[index])238 239 def __len__(self):240 return len(self.audiopaths_sid_text)241 242 243class TextAudioSpeakerCollate():244 """ Zero-pads model inputs and targets245 """246 def __init__(self, return_ids=False):247 self.return_ids = return_ids248 249 def __call__(self, batch):250 """Collate's training batch from normalized text, audio and speaker identities251 PARAMS252 ------253 batch: [text_normalized, spec_normalized, wav_normalized, sid]254 """255 # Right zero-pad all one-hot text sequences to max input length256 _, ids_sorted_decreasing = torch.sort(257 torch.LongTensor([x[1].size(1) for x in batch]),258 dim=0, descending=True)259 260 max_text_len = max([len(x[0]) for x in batch])261 max_spec_len = max([x[1].size(1) for x in batch])262 max_wav_len = max([x[2].size(1) for x in batch])263 264 text_lengths = torch.LongTensor(len(batch))265 spec_lengths = torch.LongTensor(len(batch))266 wav_lengths = torch.LongTensor(len(batch))267 sid = torch.LongTensor(len(batch))268 269 text_padded = torch.LongTensor(len(batch), max_text_len)270 spec_padded = torch.FloatTensor(len(batch), batch[0][1].size(0), max_spec_len)271 wav_padded = torch.FloatTensor(len(batch), 1, max_wav_len)272 text_padded.zero_()273 spec_padded.zero_()274 wav_padded.zero_()275 for i in range(len(ids_sorted_decreasing)):276 row = batch[ids_sorted_decreasing[i]]277 278 text = row[0]279 text_padded[i, :text.size(0)] = text280 text_lengths[i] = text.size(0)281 282 spec = row[1]283 spec_padded[i, :, :spec.size(1)] = spec284 spec_lengths[i] = spec.size(1)285 286 wav = row[2]287 wav_padded[i, :, :wav.size(1)] = wav288 wav_lengths[i] = wav.size(1)289 290 sid[i] = row[3]291 292 if self.return_ids:293 return text_padded, text_lengths, spec_padded, spec_lengths, wav_padded, wav_lengths, sid, ids_sorted_decreasing294 return text_padded, text_lengths, spec_padded, spec_lengths, wav_padded, wav_lengths, sid295 296 297class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler):298 """299 Maintain similar input lengths in a batch.300 Length groups are specified by boundaries.301 Ex) boundaries = [b1, b2, b3] -> any batch is included either {x | b1 < length(x) <=b2} or {x | b2 < length(x) <= b3}.302 303 It removes samples which are not included in the boundaries.304 Ex) boundaries = [b1, b2, b3] -> any x s.t. length(x) <= b1 or length(x) > b3 are discarded.305 """306 def __init__(self, dataset, batch_size, boundaries, num_replicas=None, rank=None, shuffle=True):307 super().__init__(dataset, num_replicas=num_replicas, rank=rank, shuffle=shuffle)308 self.lengths = dataset.lengths309 self.batch_size = batch_size310 self.boundaries = boundaries311 312 self.buckets, self.num_samples_per_bucket = self._create_buckets()313 self.total_size = sum(self.num_samples_per_bucket)314 self.num_samples = self.total_size // self.num_replicas315 316 def _create_buckets(self):317 buckets = [[] for _ in range(len(self.boundaries) - 1)]318 for i in range(len(self.lengths)):319 length = self.lengths[i]320 idx_bucket = self._bisect(length)321 if idx_bucket != -1:322 buckets[idx_bucket].append(i)323 324 for i in range(len(buckets) - 1, 0, -1):325 if len(buckets[i]) == 0:326 buckets.pop(i)327 self.boundaries.pop(i+1)328 329 num_samples_per_bucket = []330 for i in range(len(buckets)):331 len_bucket = len(buckets[i])332 total_batch_size = self.num_replicas * self.batch_size333 rem = (total_batch_size - (len_bucket % total_batch_size)) % total_batch_size334 num_samples_per_bucket.append(len_bucket + rem)335 return buckets, num_samples_per_bucket336 337 def __iter__(self):338 # deterministically shuffle based on epoch339 g = torch.Generator()340 g.manual_seed(self.epoch)341 342 indices = []343 if self.shuffle:344 for bucket in self.buckets:345 indices.append(torch.randperm(len(bucket), generator=g).tolist())346 else:347 for bucket in self.buckets:348 indices.append(list(range(len(bucket))))349 350 batches = []351 for i in range(len(self.buckets)):352 bucket = self.buckets[i]353 len_bucket = len(bucket)354 ids_bucket = indices[i]355 num_samples_bucket = self.num_samples_per_bucket[i]356 357 # add extra samples to make it evenly divisible358 rem = num_samples_bucket - len_bucket359 ids_bucket = ids_bucket + ids_bucket * (rem // len_bucket) + ids_bucket[:(rem % len_bucket)]360 361 # subsample362 ids_bucket = ids_bucket[self.rank::self.num_replicas]363 364 # batching365 for j in range(len(ids_bucket) // self.batch_size):366 batch = [bucket[idx] for idx in ids_bucket[j*self.batch_size:(j+1)*self.batch_size]]367 batches.append(batch)368 369 if self.shuffle:370 batch_ids = torch.randperm(len(batches), generator=g).tolist()371 batches = [batches[i] for i in batch_ids]372 self.batches = batches373 374 assert len(self.batches) * self.batch_size == self.num_samples375 return iter(self.batches)376 377 def _bisect(self, x, lo=0, hi=None):378 if hi is None:379 hi = len(self.boundaries) - 1380 381 if hi > lo:382 mid = (hi + lo) // 2383 if self.boundaries[mid] < x and x <= self.boundaries[mid+1]:384 return mid385 elif x <= self.boundaries[mid]:386 return self._bisect(x, lo, mid)387 else:388 return self._bisect(x, mid + 1, hi)389 else:390 return -1391 392 def __len__(self):393 return self.num_samples // self.batch_size394 