mosi77/RVC_HFv2
0
1import os, traceback2import numpy as np3import torch4import torch.utils.data5 6from mel_processing import spectrogram_torch7from utils import load_wav_to_torch, load_filepaths_and_text8 9 10class TextAudioLoaderMultiNSFsid(torch.utils.data.Dataset):11 """12 1) loads audio, text pairs13 2) normalizes text and converts them to sequences of integers14 3) computes spectrograms from audio files.15 """16 17 def __init__(self, audiopaths_and_text, hparams):18 self.audiopaths_and_text = load_filepaths_and_text(audiopaths_and_text)19 self.max_wav_value = hparams.max_wav_value20 self.sampling_rate = hparams.sampling_rate21 self.filter_length = hparams.filter_length22 self.hop_length = hparams.hop_length23 self.win_length = hparams.win_length24 self.sampling_rate = hparams.sampling_rate25 self.min_text_len = getattr(hparams, "min_text_len", 1)26 self.max_text_len = getattr(hparams, "max_text_len", 5000)27 self._filter()28 29 def _filter(self):30 """31 Filter text & store spec lengths32 """33 # Store spectrogram lengths for Bucketing34 # wav_length ~= file_size / (wav_channels * Bytes per dim) = file_size / (1 * 2)35 # spec_length = wav_length // hop_length36 audiopaths_and_text_new = []37 lengths = []38 for audiopath, text, pitch, pitchf, dv in self.audiopaths_and_text:39 if self.min_text_len <= len(text) and len(text) <= self.max_text_len:40 audiopaths_and_text_new.append([audiopath, text, pitch, pitchf, dv])41 lengths.append(os.path.getsize(audiopath) // (3 * self.hop_length))42 self.audiopaths_and_text = audiopaths_and_text_new43 self.lengths = lengths44 45 def get_sid(self, sid):46 sid = torch.LongTensor([int(sid)])47 return sid48 49 def get_audio_text_pair(self, audiopath_and_text):50 # separate filename and text51 file = audiopath_and_text[0]52 phone = audiopath_and_text[1]53 pitch = audiopath_and_text[2]54 pitchf = audiopath_and_text[3]55 dv = audiopath_and_text[4]56 57 phone, pitch, pitchf = self.get_labels(phone, pitch, pitchf)58 spec, wav = self.get_audio(file)59 dv = self.get_sid(dv)60 61 len_phone = phone.size()[0]62 len_spec = spec.size()[-1]63 # print(123,phone.shape,pitch.shape,spec.shape)64 if len_phone != len_spec:65 len_min = min(len_phone, len_spec)66 # amor67 len_wav = len_min * self.hop_length68 69 spec = spec[:, :len_min]70 wav = wav[:, :len_wav]71 72 phone = phone[:len_min, :]73 pitch = pitch[:len_min]74 pitchf = pitchf[:len_min]75 76 return (spec, wav, phone, pitch, pitchf, dv)77 78 def get_labels(self, phone, pitch, pitchf):79 phone = np.load(phone)80 phone = np.repeat(phone, 2, axis=0)81 pitch = np.load(pitch)82 pitchf = np.load(pitchf)83 n_num = min(phone.shape[0], 900) # DistributedBucketSampler84 # print(234,phone.shape,pitch.shape)85 phone = phone[:n_num, :]86 pitch = pitch[:n_num]87 pitchf = pitchf[:n_num]88 phone = torch.FloatTensor(phone)89 pitch = torch.LongTensor(pitch)90 pitchf = torch.FloatTensor(pitchf)91 return phone, pitch, pitchf92 93 def get_audio(self, filename):94 audio, sampling_rate = load_wav_to_torch(filename)95 if sampling_rate != self.sampling_rate:96 raise ValueError(97 "{} SR doesn't match target {} SR".format(98 sampling_rate, self.sampling_rate99 )100 )101 audio_norm = audio102 # audio_norm = audio / self.max_wav_value103 # audio_norm = audio / np.abs(audio).max()104 105 audio_norm = audio_norm.unsqueeze(0)106 spec_filename = filename.replace(".wav", ".spec.pt")107 if os.path.exists(spec_filename):108 try:109 spec = torch.load(spec_filename)110 except:111 print(spec_filename, traceback.format_exc())112 spec = spectrogram_torch(113 audio_norm,114 self.filter_length,115 self.sampling_rate,116 self.hop_length,117 self.win_length,118 center=False,119 )120 spec = torch.squeeze(spec, 0)121 torch.save(spec, spec_filename, _use_new_zipfile_serialization=False)122 else:123 spec = spectrogram_torch(124 audio_norm,125 self.filter_length,126 self.sampling_rate,127 self.hop_length,128 self.win_length,129 center=False,130 )131 spec = torch.squeeze(spec, 0)132 torch.save(spec, spec_filename, _use_new_zipfile_serialization=False)133 return spec, audio_norm134 135 def __getitem__(self, index):136 return self.get_audio_text_pair(self.audiopaths_and_text[index])137 138 def __len__(self):139 return len(self.audiopaths_and_text)140 141 142class TextAudioCollateMultiNSFsid:143 """Zero-pads model inputs and targets"""144 145 def __init__(self, return_ids=False):146 self.return_ids = return_ids147 148 def __call__(self, batch):149 """Collate's training batch from normalized text and aduio150 PARAMS151 ------152 batch: [text_normalized, spec_normalized, wav_normalized]153 """154 # Right zero-pad all one-hot text sequences to max input length155 _, ids_sorted_decreasing = torch.sort(156 torch.LongTensor([x[0].size(1) for x in batch]), dim=0, descending=True157 )158 159 max_spec_len = max([x[0].size(1) for x in batch])160 max_wave_len = max([x[1].size(1) for x in batch])161 spec_lengths = torch.LongTensor(len(batch))162 wave_lengths = torch.LongTensor(len(batch))163 spec_padded = torch.FloatTensor(len(batch), batch[0][0].size(0), max_spec_len)164 wave_padded = torch.FloatTensor(len(batch), 1, max_wave_len)165 spec_padded.zero_()166 wave_padded.zero_()167 168 max_phone_len = max([x[2].size(0) for x in batch])169 phone_lengths = torch.LongTensor(len(batch))170 phone_padded = torch.FloatTensor(171 len(batch), max_phone_len, batch[0][2].shape[1]172 ) # (spec, wav, phone, pitch)173 pitch_padded = torch.LongTensor(len(batch), max_phone_len)174 pitchf_padded = torch.FloatTensor(len(batch), max_phone_len)175 phone_padded.zero_()176 pitch_padded.zero_()177 pitchf_padded.zero_()178 # dv = torch.FloatTensor(len(batch), 256)#gin=256179 sid = torch.LongTensor(len(batch))180 181 for i in range(len(ids_sorted_decreasing)):182 row = batch[ids_sorted_decreasing[i]]183 184 spec = row[0]185 spec_padded[i, :, : spec.size(1)] = spec186 spec_lengths[i] = spec.size(1)187 188 wave = row[1]189 wave_padded[i, :, : wave.size(1)] = wave190 wave_lengths[i] = wave.size(1)191 192 phone = row[2]193 phone_padded[i, : phone.size(0), :] = phone194 phone_lengths[i] = phone.size(0)195 196 pitch = row[3]197 pitch_padded[i, : pitch.size(0)] = pitch198 pitchf = row[4]199 pitchf_padded[i, : pitchf.size(0)] = pitchf200 201 # dv[i] = row[5]202 sid[i] = row[5]203 204 return (205 phone_padded,206 phone_lengths,207 pitch_padded,208 pitchf_padded,209 spec_padded,210 spec_lengths,211 wave_padded,212 wave_lengths,213 # dv214 sid,215 )216 217 218class TextAudioLoader(torch.utils.data.Dataset):219 """220 1) loads audio, text pairs221 2) normalizes text and converts them to sequences of integers222 3) computes spectrograms from audio files.223 """224 225 def __init__(self, audiopaths_and_text, hparams):226 self.audiopaths_and_text = load_filepaths_and_text(audiopaths_and_text)227 self.max_wav_value = hparams.max_wav_value228 self.sampling_rate = hparams.sampling_rate229 self.filter_length = hparams.filter_length230 self.hop_length = hparams.hop_length231 self.win_length = hparams.win_length232 self.sampling_rate = hparams.sampling_rate233 self.min_text_len = getattr(hparams, "min_text_len", 1)234 self.max_text_len = getattr(hparams, "max_text_len", 5000)235 self._filter()236 237 def _filter(self):238 """239 Filter text & store spec lengths240 """241 # Store spectrogram lengths for Bucketing242 # wav_length ~= file_size / (wav_channels * Bytes per dim) = file_size / (1 * 2)243 # spec_length = wav_length // hop_length244 audiopaths_and_text_new = []245 lengths = []246 for audiopath, text, dv in self.audiopaths_and_text:247 if self.min_text_len <= len(text) and len(text) <= self.max_text_len:248 audiopaths_and_text_new.append([audiopath, text, dv])249 lengths.append(os.path.getsize(audiopath) // (3 * self.hop_length))250 self.audiopaths_and_text = audiopaths_and_text_new251 self.lengths = lengths252 253 def get_sid(self, sid):254 sid = torch.LongTensor([int(sid)])255 return sid256 257 def get_audio_text_pair(self, audiopath_and_text):258 # separate filename and text259 file = audiopath_and_text[0]260 phone = audiopath_and_text[1]261 dv = audiopath_and_text[2]262 263 phone = self.get_labels(phone)264 spec, wav = self.get_audio(file)265 dv = self.get_sid(dv)266 267 len_phone = phone.size()[0]268 len_spec = spec.size()[-1]269 if len_phone != len_spec:270 len_min = min(len_phone, len_spec)271 len_wav = len_min * self.hop_length272 spec = spec[:, :len_min]273 wav = wav[:, :len_wav]274 phone = phone[:len_min, :]275 return (spec, wav, phone, dv)276 277 def get_labels(self, phone):278 phone = np.load(phone)279 phone = np.repeat(phone, 2, axis=0)280 n_num = min(phone.shape[0], 900) # DistributedBucketSampler281 phone = phone[:n_num, :]282 phone = torch.FloatTensor(phone)283 return phone284 285 def get_audio(self, filename):286 audio, sampling_rate = load_wav_to_torch(filename)287 if sampling_rate != self.sampling_rate:288 raise ValueError(289 "{} SR doesn't match target {} SR".format(290 sampling_rate, self.sampling_rate291 )292 )293 audio_norm = audio294 # audio_norm = audio / self.max_wav_value295 # audio_norm = audio / np.abs(audio).max()296 297 audio_norm = audio_norm.unsqueeze(0)298 spec_filename = filename.replace(".wav", ".spec.pt")299 if os.path.exists(spec_filename):300 try:301 spec = torch.load(spec_filename)302 except:303 print(spec_filename, traceback.format_exc())304 spec = spectrogram_torch(305 audio_norm,306 self.filter_length,307 self.sampling_rate,308 self.hop_length,309 self.win_length,310 center=False,311 )312 spec = torch.squeeze(spec, 0)313 torch.save(spec, spec_filename, _use_new_zipfile_serialization=False)314 else:315 spec = spectrogram_torch(316 audio_norm,317 self.filter_length,318 self.sampling_rate,319 self.hop_length,320 self.win_length,321 center=False,322 )323 spec = torch.squeeze(spec, 0)324 torch.save(spec, spec_filename, _use_new_zipfile_serialization=False)325 return spec, audio_norm326 327 def __getitem__(self, index):328 return self.get_audio_text_pair(self.audiopaths_and_text[index])329 330 def __len__(self):331 return len(self.audiopaths_and_text)332 333 334class TextAudioCollate:335 """Zero-pads model inputs and targets"""336 337 def __init__(self, return_ids=False):338 self.return_ids = return_ids339 340 def __call__(self, batch):341 """Collate's training batch from normalized text and aduio342 PARAMS343 ------344 batch: [text_normalized, spec_normalized, wav_normalized]345 """346 # Right zero-pad all one-hot text sequences to max input length347 _, ids_sorted_decreasing = torch.sort(348 torch.LongTensor([x[0].size(1) for x in batch]), dim=0, descending=True349 )350 351 max_spec_len = max([x[0].size(1) for x in batch])352 max_wave_len = max([x[1].size(1) for x in batch])353 spec_lengths = torch.LongTensor(len(batch))354 wave_lengths = torch.LongTensor(len(batch))355 spec_padded = torch.FloatTensor(len(batch), batch[0][0].size(0), max_spec_len)356 wave_padded = torch.FloatTensor(len(batch), 1, max_wave_len)357 spec_padded.zero_()358 wave_padded.zero_()359 360 max_phone_len = max([x[2].size(0) for x in batch])361 phone_lengths = torch.LongTensor(len(batch))362 phone_padded = torch.FloatTensor(363 len(batch), max_phone_len, batch[0][2].shape[1]364 )365 phone_padded.zero_()366 sid = torch.LongTensor(len(batch))367 368 for i in range(len(ids_sorted_decreasing)):369 row = batch[ids_sorted_decreasing[i]]370 371 spec = row[0]372 spec_padded[i, :, : spec.size(1)] = spec373 spec_lengths[i] = spec.size(1)374 375 wave = row[1]376 wave_padded[i, :, : wave.size(1)] = wave377 wave_lengths[i] = wave.size(1)378 379 phone = row[2]380 phone_padded[i, : phone.size(0), :] = phone381 phone_lengths[i] = phone.size(0)382 383 sid[i] = row[3]384 385 return (386 phone_padded,387 phone_lengths,388 spec_padded,389 spec_lengths,390 wave_padded,391 wave_lengths,392 sid,393 )394 395 396class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler):397 """398 Maintain similar input lengths in a batch.399 Length groups are specified by boundaries.400 Ex) boundaries = [b1, b2, b3] -> any batch is included either {x | b1 < length(x) <=b2} or {x | b2 < length(x) <= b3}.401 402 It removes samples which are not included in the boundaries.403 Ex) boundaries = [b1, b2, b3] -> any x s.t. length(x) <= b1 or length(x) > b3 are discarded.404 """405 406 def __init__(407 self,408 dataset,409 batch_size,410 boundaries,411 num_replicas=None,412 rank=None,413 shuffle=True,414 ):415 super().__init__(dataset, num_replicas=num_replicas, rank=rank, shuffle=shuffle)416 self.lengths = dataset.lengths417 self.batch_size = batch_size418 self.boundaries = boundaries419 420 self.buckets, self.num_samples_per_bucket = self._create_buckets()421 self.total_size = sum(self.num_samples_per_bucket)422 self.num_samples = self.total_size // self.num_replicas423 424 def _create_buckets(self):425 buckets = [[] for _ in range(len(self.boundaries) - 1)]426 for i in range(len(self.lengths)):427 length = self.lengths[i]428 idx_bucket = self._bisect(length)429 if idx_bucket != -1:430 buckets[idx_bucket].append(i)431 432 for i in range(len(buckets) - 1, -1, -1): #433 if len(buckets[i]) == 0:434 buckets.pop(i)435 self.boundaries.pop(i + 1)436 437 num_samples_per_bucket = []438 for i in range(len(buckets)):439 len_bucket = len(buckets[i])440 total_batch_size = self.num_replicas * self.batch_size441 rem = (442 total_batch_size - (len_bucket % total_batch_size)443 ) % total_batch_size444 num_samples_per_bucket.append(len_bucket + rem)445 return buckets, num_samples_per_bucket446 447 def __iter__(self):448 # deterministically shuffle based on epoch449 g = torch.Generator()450 g.manual_seed(self.epoch)451 452 indices = []453 if self.shuffle:454 for bucket in self.buckets:455 indices.append(torch.randperm(len(bucket), generator=g).tolist())456 else:457 for bucket in self.buckets:458 indices.append(list(range(len(bucket))))459 460 batches = []461 for i in range(len(self.buckets)):462 bucket = self.buckets[i]463 len_bucket = len(bucket)464 ids_bucket = indices[i]465 num_samples_bucket = self.num_samples_per_bucket[i]466 467 # add extra samples to make it evenly divisible468 rem = num_samples_bucket - len_bucket469 ids_bucket = (470 ids_bucket471 + ids_bucket * (rem // len_bucket)472 + ids_bucket[: (rem % len_bucket)]473 )474 475 # subsample476 ids_bucket = ids_bucket[self.rank :: self.num_replicas]477 478 # batching479 for j in range(len(ids_bucket) // self.batch_size):480 batch = [481 bucket[idx]482 for idx in ids_bucket[483 j * self.batch_size : (j + 1) * self.batch_size484 ]485 ]486 batches.append(batch)487 488 if self.shuffle:489 batch_ids = torch.randperm(len(batches), generator=g).tolist()490 batches = [batches[i] for i in batch_ids]491 self.batches = batches492 493 assert len(self.batches) * self.batch_size == self.num_samples494 return iter(self.batches)495 496 def _bisect(self, x, lo=0, hi=None):497 if hi is None:498 hi = len(self.boundaries) - 1499 500 if hi > lo:501 mid = (hi + lo) // 2502 if self.boundaries[mid] < x and x <= self.boundaries[mid + 1]:503 return mid504 elif x <= self.boundaries[mid]:505 return self._bisect(x, lo, mid)506 else:507 return self._bisect(x, mid + 1, hi)508 else:509 return -1510 511 def __len__(self):512 return self.num_samples // self.batch_size513 