CoolFace
Apppublic

Aloento/9Nine-PITS

sourceHugging Faceagpl-3.0updated 4y agoView on Hugging Face
1likes
data_utils.py359 linesDownload Raw Back to root
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