hanfish/LSai
0
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 