CoolFace
Apppublic

hanfish/LSai

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
data_utils.py327 linesDownload Raw Back to module
1import time,logging2import os3import random,traceback4import numpy as np5import torch6import torch.utils.data7from tqdm import tqdm8 9from module import commons10from module.mel_processing import spectrogram_torch11from text import cleaned_text_to_sequence12from utils import load_wav_to_torch, load_filepaths_and_text13import torch.nn.functional as F14from functools import lru_cache15import torch16import requests17from scipy.io import wavfile18from io import BytesIO19# from config import exp_dir20from my_utils import load_audio21 22class TextAudioSpeakerLoader(torch.utils.data.Dataset):23    """24        1) loads audio, speaker_id, text pairs25        2) normalizes text and converts them to sequences of integers26        3) computes spectrograms from audio files.27    """28 29    def __init__(self, hparams, val=False):30        exp_dir=hparams.exp_dir31        self.path2="%s/2-name2text.txt"%exp_dir32        self.path4="%s/4-cnhubert"%exp_dir33        self.path5="%s/5-wav32k"%exp_dir34        assert os.path.exists(self.path2)35        assert os.path.exists(self.path4)36        assert os.path.exists(self.path5)37        names4=set([name[:-3]for name in list(os.listdir(self.path4))])#去除.pt后缀38        names5=set(os.listdir(self.path5))39        self.phoneme_data={}40        with open(self.path2,"r",encoding="utf8")as f:41            lines=f.read().strip("\n").split("\n")42 43        for line in lines:44            tmp=line.split("\t")45            if(len(tmp)!=4):continue46            self.phoneme_data[tmp[0]]=[tmp[1]]47 48        self.audiopaths_sid_text=list(set(self.phoneme_data)&names4&names5)49        tmp=self.audiopaths_sid_text50        leng=len(tmp)51        min_num=10052        if(leng<min_num):53            self.audiopaths_sid_text=[]54            for _ in range(max(2, int(min_num / leng))):55                self.audiopaths_sid_text += tmp56        self.max_wav_value = hparams.max_wav_value57        self.sampling_rate = hparams.sampling_rate58        self.filter_length = hparams.filter_length59        self.hop_length = hparams.hop_length60        self.win_length = hparams.win_length61        self.sampling_rate = hparams.sampling_rate62        self.val = val63 64        random.seed(1234)65        random.shuffle(self.audiopaths_sid_text)66 67        print("phoneme_data_len:", len(self.phoneme_data.keys()))68        print("wav_data_len:", len(self.audiopaths_sid_text))69 70        audiopaths_sid_text_new = []71        lengths = []72        skipped_phone  = 073        skipped_dur   = 074        for audiopath in tqdm(self.audiopaths_sid_text):75            try:76                phoneme = self.phoneme_data[audiopath][0]77                phoneme = phoneme.split(' ')78                phoneme_ids = cleaned_text_to_sequence(phoneme)79            except Exception:80                print(f"{audiopath} not in self.phoneme_data !")81                skipped_phone  += 182                continue83            size=os.path.getsize("%s/%s"%(self.path5,audiopath))84            duration = size / self.sampling_rate / 285            if (54 > duration > 0.6 or self.val):86                audiopaths_sid_text_new.append([audiopath, phoneme_ids])87                lengths.append(size // (2 * self.hop_length))88            else:89                skipped_dur += 190                continue91        print("skipped_phone: ", skipped_phone, ", skipped_dur: ", skipped_dur)92        print("total left: ", len(audiopaths_sid_text_new))93        assert len(audiopaths_sid_text_new)>1#至少能凑够batch size,这里todo94        self.audiopaths_sid_text = audiopaths_sid_text_new95        self.lengths = lengths96 97    def get_audio_text_speaker_pair(self, audiopath_sid_text):98        audiopath, phoneme_ids = audiopath_sid_text99        text = torch.FloatTensor(phoneme_ids)100        try:101            spec, wav = self.get_audio("%s/%s"%(self.path5,audiopath))102            with torch.no_grad():103                ssl = torch.load("%s/%s.pt"%(self.path4,audiopath),map_location="cpu")104                if(ssl.shape[-1]!=spec.shape[-1]):105                    typee=ssl.dtype106                    ssl=F.pad(ssl.float(),(0,1),mode="replicate").to(typee)107                ssl.requires_grad=False108        except:109            traceback.print_exc()110            spec = torch.zeros(1025, 100)111            wav = torch.zeros(1, 100*self.hop_length)112            ssl=torch.zeros(1,768,100)113            text=text[-1:]114            print("load audio or ssl error!!!!!!", audiopath)115        # print(ssl.requires_grad,spec.requires_grad,wav.requires_grad,text.requires_grad)116        return (ssl, spec, wav, text)117 118    def get_audio(self, filename):119        audio_array = load_audio(filename,self.sampling_rate)#load_audio的方法是已经归一化到-1~1之间的,不用再/32768120        # print(filename,audio_array.max(),audio_array.min(),audio_array.mean())121        audio=torch.FloatTensor(audio_array)#/32768122        audio_norm = audio123        audio_norm = audio_norm.unsqueeze(0)124        spec = spectrogram_torch(audio_norm, self.filter_length,self.sampling_rate, self.hop_length, self.win_length,center=False)125        spec = torch.squeeze(spec, 0)126        return spec, audio_norm127 128    def get_sid(self, sid):129        sid = torch.LongTensor([int(sid)])130        return sid131 132    def __getitem__(self, index):133        # with torch.no_grad():134            return self.get_audio_text_speaker_pair(self.audiopaths_sid_text[index])135 136    def __len__(self):137        return len(self.audiopaths_sid_text)138 139    def random_slice(self, ssl, wav, mel):140        assert abs(ssl.shape[-1]- wav.shape[-1]//self.hop_length) < 3, ("first", ssl.shape, wav.shape)141 142        len_mel = mel.shape[1]143        if self.val:144            reference_mel = mel[:, :len_mel//3]145            return reference_mel, ssl, wav, mel146        dir = random.randint(0, 1)147        sep_point = random.randint(int(len_mel//3), int(len_mel//3*2))148 149        if dir == 0:150            reference_mel = mel[:, :sep_point]151            ssl = ssl[:, :, sep_point:]152            wav2 = wav[:, sep_point*self.hop_length:]153            mel = mel[:, sep_point:]154        else:155            reference_mel = mel[:, sep_point:]156            ssl = ssl[:, :, :sep_point]157            wav2 = wav[:, :sep_point*self.hop_length]158            mel = mel[:, :sep_point]159 160        assert abs(ssl.shape[-1]- wav2.shape[-1]//self.hop_length) < 3, (ssl.shape, wav.shape,wav2.shape, mel.shape, sep_point,self.hop_length, sep_point*self.hop_length, dir)161        return reference_mel, ssl, wav2, mel162 163 164class TextAudioSpeakerCollate():165    """ Zero-pads model inputs and targets166    """167 168    def __init__(self, return_ids=False):169        self.return_ids = return_ids170 171    def __call__(self, batch):172        """Collate's training batch from normalized text, audio and speaker identities173        PARAMS174        ------175        batch: [text_normalized, spec_normalized, wav_normalized, sid]176        """177        # Right zero-pad all one-hot text sequences to max input length178        _, ids_sorted_decreasing = torch.sort(179            torch.LongTensor([x[1].size(1) for x in batch]),180            dim=0, descending=True)181 182        max_ssl_len = max([x[0].size(2) for x in batch])183        max_ssl_len = int(2 * ((max_ssl_len // 2) + 1))184        max_spec_len = max([x[1].size(1) for x in batch])185        max_spec_len = int(2 * ((max_spec_len // 2) + 1))186        max_wav_len = max([x[2].size(1) for x in batch])187        max_text_len = max([x[3].size(0) for x in batch])188 189        ssl_lengths = torch.LongTensor(len(batch))190        spec_lengths = torch.LongTensor(len(batch))191        wav_lengths = torch.LongTensor(len(batch))192        text_lengths = torch.LongTensor(len(batch))193 194        spec_padded = torch.FloatTensor(len(batch), batch[0][1].size(0), max_spec_len)195        wav_padded = torch.FloatTensor(len(batch), 1, max_wav_len)196        ssl_padded = torch.FloatTensor(len(batch), batch[0][0].size(1), max_ssl_len)197        text_padded = torch.LongTensor(len(batch),  max_text_len)198 199        spec_padded.zero_()200        wav_padded.zero_()201        ssl_padded.zero_()202        text_padded.zero_()203 204        for i in range(len(ids_sorted_decreasing)):205            row = batch[ids_sorted_decreasing[i]]206 207            ssl = row[0]208            ssl_padded[i, :, :ssl.size(2)] = ssl[0, :, :]209            ssl_lengths[i] = ssl.size(2)210 211            spec = row[1]212            spec_padded[i, :, :spec.size(1)] = spec213            spec_lengths[i] = spec.size(1)214 215            wav = row[2]216            wav_padded[i, :, :wav.size(1)] = wav217            wav_lengths[i] = wav.size(1)218 219            text = row[3]220            text_padded[i, :text.size(0)] = text221            text_lengths[i] = text.size(0)222 223 224        return ssl_padded, ssl_lengths, spec_padded, spec_lengths, wav_padded, wav_lengths, text_padded, text_lengths225 226 227class DistributedBucketSampler(torch.utils.data.distributed.DistributedSampler):228    """229    Maintain similar input lengths in a batch.230    Length groups are specified by boundaries.231    Ex) boundaries = [b1, b2, b3] -> any batch is included either {x | b1 < length(x) <=b2} or {x | b2 < length(x) <= b3}.232 233    It removes samples which are not included in the boundaries.234    Ex) boundaries = [b1, b2, b3] -> any x s.t. length(x) <= b1 or length(x) > b3 are discarded.235    """236 237    def __init__(self, dataset, batch_size, boundaries, num_replicas=None, rank=None, shuffle=True):238        super().__init__(dataset, num_replicas=num_replicas, rank=rank, shuffle=shuffle)239        self.lengths = dataset.lengths240        # print(233333333333333,self.lengths,dir(dataset))241        self.batch_size = batch_size242        self.boundaries = boundaries243 244        self.buckets, self.num_samples_per_bucket = self._create_buckets()245        self.total_size = sum(self.num_samples_per_bucket)246        self.num_samples = self.total_size // self.num_replicas247 248    def _create_buckets(self):249        buckets = [[] for _ in range(len(self.boundaries) - 1)]250        for i in range(len(self.lengths)):251            length = self.lengths[i]252            idx_bucket = self._bisect(length)253            if idx_bucket != -1:254                buckets[idx_bucket].append(i)255 256        for i in range(len(buckets) - 1, 0, -1):257        # for i in range(len(buckets) - 1, -1, -1):258            if len(buckets[i]) == 0:259                buckets.pop(i)260                self.boundaries.pop(i + 1)261 262        num_samples_per_bucket = []263        for i in range(len(buckets)):264            len_bucket = len(buckets[i])265            total_batch_size = self.num_replicas * self.batch_size266            rem = (total_batch_size - (len_bucket % total_batch_size)) % total_batch_size267            num_samples_per_bucket.append(len_bucket + rem)268        return buckets, num_samples_per_bucket269 270    def __iter__(self):271        # deterministically shuffle based on epoch272        g = torch.Generator()273        g.manual_seed(self.epoch)274 275        indices = []276        if self.shuffle:277            for bucket in self.buckets:278                indices.append(torch.randperm(len(bucket), generator=g).tolist())279        else:280            for bucket in self.buckets:281                indices.append(list(range(len(bucket))))282 283        batches = []284        for i in range(len(self.buckets)):285            bucket = self.buckets[i]286            len_bucket = len(bucket)287            ids_bucket = indices[i]288            num_samples_bucket = self.num_samples_per_bucket[i]289 290            # add extra samples to make it evenly divisible291            rem = num_samples_bucket - len_bucket292            ids_bucket = ids_bucket + ids_bucket * (rem // len_bucket) + ids_bucket[:(rem % len_bucket)]293 294            # subsample295            ids_bucket = ids_bucket[self.rank::self.num_replicas]296 297            # batching298            for j in range(len(ids_bucket) // self.batch_size):299                batch = [bucket[idx] for idx in ids_bucket[j * self.batch_size:(j + 1) * self.batch_size]]300                batches.append(batch)301 302        if self.shuffle:303            batch_ids = torch.randperm(len(batches), generator=g).tolist()304            batches = [batches[i] for i in batch_ids]305        self.batches = batches306 307        assert len(self.batches) * self.batch_size == self.num_samples308        return iter(self.batches)309 310    def _bisect(self, x, lo=0, hi=None):311        if hi is None:312            hi = len(self.boundaries) - 1313 314        if hi > lo:315            mid = (hi + lo) // 2316            if self.boundaries[mid] < x and x <= self.boundaries[mid + 1]:317                return mid318            elif x <= self.boundaries[mid]:319                return self._bisect(x, lo, mid)320            else:321                return self._bisect(x, mid + 1, hi)322        else:323            return -1324 325    def __len__(self):326        return self.num_samples // self.batch_size327