CoolFace
Apppublic

mosi77/RVC_HFv2

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
data_utils.py513 linesDownload Raw Back to train
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