Team Ai
Apppublic

Kafke/Code-Realize-TTS

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
data_utils.py394 linesDownload Raw Back to root
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