justyoung/DiffSinger
1
1import matplotlib2 3matplotlib.use('Agg')4 5from utils import audio6import matplotlib.pyplot as plt7from data_gen.tts.data_gen_utils import get_pitch8from tasks.tts.fs2_utils import FastSpeechDataset9from utils.cwt import cwt2f010from utils.pl_utils import data_loader11import os12from multiprocessing.pool import Pool13from tqdm import tqdm14from modules.fastspeech.tts_modules import mel2ph_to_dur15from utils.hparams import hparams16from utils.plot import spec_to_figure, dur_to_figure, f0_to_figure17from utils.pitch_utils import denorm_f018from modules.fastspeech.fs2 import FastSpeech219from tasks.tts.tts import TtsTask20import torch21import torch.optim22import torch.utils.data23import torch.nn.functional as F24import utils25import torch.distributions26import numpy as np27from modules.commons.ssim import ssim28 29class FastSpeech2Task(TtsTask):30 def __init__(self):31 super(FastSpeech2Task, self).__init__()32 self.dataset_cls = FastSpeechDataset33 self.mse_loss_fn = torch.nn.MSELoss()34 mel_losses = hparams['mel_loss'].split("|")35 self.loss_and_lambda = {}36 for i, l in enumerate(mel_losses):37 if l == '':38 continue39 if ':' in l:40 l, lbd = l.split(":")41 lbd = float(lbd)42 else:43 lbd = 1.044 self.loss_and_lambda[l] = lbd45 print("| Mel losses:", self.loss_and_lambda)46 self.sil_ph = self.phone_encoder.sil_phonemes()47 48 @data_loader49 def train_dataloader(self):50 train_dataset = self.dataset_cls(hparams['train_set_name'], shuffle=True)51 return self.build_dataloader(train_dataset, True, self.max_tokens, self.max_sentences,52 endless=hparams['endless_ds'])53 54 @data_loader55 def val_dataloader(self):56 valid_dataset = self.dataset_cls(hparams['valid_set_name'], shuffle=False)57 return self.build_dataloader(valid_dataset, False, self.max_eval_tokens, self.max_eval_sentences)58 59 @data_loader60 def test_dataloader(self):61 test_dataset = self.dataset_cls(hparams['test_set_name'], shuffle=False)62 return self.build_dataloader(test_dataset, False, self.max_eval_tokens,63 self.max_eval_sentences, batch_by_size=False)64 65 def build_tts_model(self):66 self.model = FastSpeech2(self.phone_encoder)67 68 def build_model(self):69 self.build_tts_model()70 if hparams['load_ckpt'] != '':71 self.load_ckpt(hparams['load_ckpt'], strict=True)72 utils.print_arch(self.model)73 return self.model74 75 def _training_step(self, sample, batch_idx, _):76 loss_output = self.run_model(self.model, sample)77 total_loss = sum([v for v in loss_output.values() if isinstance(v, torch.Tensor) and v.requires_grad])78 loss_output['batch_size'] = sample['txt_tokens'].size()[0]79 return total_loss, loss_output80 81 def validation_step(self, sample, batch_idx):82 outputs = {}83 outputs['losses'] = {}84 outputs['losses'], model_out = self.run_model(self.model, sample, return_output=True)85 outputs['total_loss'] = sum(outputs['losses'].values())86 outputs['nsamples'] = sample['nsamples']87 mel_out = self.model.out2mel(model_out['mel_out'])88 outputs = utils.tensors_to_scalars(outputs)89 # if sample['mels'].shape[0] == 1:90 # self.add_laplace_var(mel_out, sample['mels'], outputs)91 if batch_idx < hparams['num_valid_plots']:92 self.plot_mel(batch_idx, sample['mels'], mel_out)93 self.plot_dur(batch_idx, sample, model_out)94 if hparams['use_pitch_embed']:95 self.plot_pitch(batch_idx, sample, model_out)96 return outputs97 98 def _validation_end(self, outputs):99 all_losses_meter = {100 'total_loss': utils.AvgrageMeter(),101 }102 for output in outputs:103 n = output['nsamples']104 for k, v in output['losses'].items():105 if k not in all_losses_meter:106 all_losses_meter[k] = utils.AvgrageMeter()107 all_losses_meter[k].update(v, n)108 all_losses_meter['total_loss'].update(output['total_loss'], n)109 return {k: round(v.avg, 4) for k, v in all_losses_meter.items()}110 111 def run_model(self, model, sample, return_output=False):112 txt_tokens = sample['txt_tokens'] # [B, T_t]113 target = sample['mels'] # [B, T_s, 80]114 mel2ph = sample['mel2ph'] # [B, T_s]115 f0 = sample['f0']116 uv = sample['uv']117 energy = sample['energy']118 spk_embed = sample.get('spk_embed') if not hparams['use_spk_id'] else sample.get('spk_ids')119 if hparams['pitch_type'] == 'cwt':120 cwt_spec = sample[f'cwt_spec']121 f0_mean = sample['f0_mean']122 f0_std = sample['f0_std']123 sample['f0_cwt'] = f0 = model.cwt2f0_norm(cwt_spec, f0_mean, f0_std, mel2ph)124 125 output = model(txt_tokens, mel2ph=mel2ph, spk_embed=spk_embed,126 ref_mels=target, f0=f0, uv=uv, energy=energy, infer=False)127 128 losses = {}129 self.add_mel_loss(output['mel_out'], target, losses)130 self.add_dur_loss(output['dur'], mel2ph, txt_tokens, losses=losses)131 if hparams['use_pitch_embed']:132 self.add_pitch_loss(output, sample, losses)133 if hparams['use_energy_embed']:134 self.add_energy_loss(output['energy_pred'], energy, losses)135 if not return_output:136 return losses137 else:138 return losses, output139 140 ############141 # losses142 ############143 def add_mel_loss(self, mel_out, target, losses, postfix='', mel_mix_loss=None):144 if mel_mix_loss is None:145 for loss_name, lbd in self.loss_and_lambda.items():146 if 'l1' == loss_name:147 l = self.l1_loss(mel_out, target)148 elif 'mse' == loss_name:149 raise NotImplementedError150 elif 'ssim' == loss_name:151 l = self.ssim_loss(mel_out, target)152 elif 'gdl' == loss_name:153 raise NotImplementedError154 losses[f'{loss_name}{postfix}'] = l * lbd155 else:156 raise NotImplementedError157 158 def l1_loss(self, decoder_output, target):159 # decoder_output : B x T x n_mel160 # target : B x T x n_mel161 l1_loss = F.l1_loss(decoder_output, target, reduction='none')162 weights = self.weights_nonzero_speech(target)163 l1_loss = (l1_loss * weights).sum() / weights.sum()164 return l1_loss165 166 def ssim_loss(self, decoder_output, target, bias=6.0):167 # decoder_output : B x T x n_mel168 # target : B x T x n_mel169 assert decoder_output.shape == target.shape170 weights = self.weights_nonzero_speech(target)171 decoder_output = decoder_output[:, None] + bias172 target = target[:, None] + bias173 ssim_loss = 1 - ssim(decoder_output, target, size_average=False)174 ssim_loss = (ssim_loss * weights).sum() / weights.sum()175 return ssim_loss176 177 def add_dur_loss(self, dur_pred, mel2ph, txt_tokens, losses=None):178 """179 180 :param dur_pred: [B, T], float, log scale181 :param mel2ph: [B, T]182 :param txt_tokens: [B, T]183 :param losses:184 :return:185 """186 B, T = txt_tokens.shape187 nonpadding = (txt_tokens != 0).float()188 dur_gt = mel2ph_to_dur(mel2ph, T).float() * nonpadding189 is_sil = torch.zeros_like(txt_tokens).bool()190 for p in self.sil_ph:191 is_sil = is_sil | (txt_tokens == self.phone_encoder.encode(p)[0])192 is_sil = is_sil.float() # [B, T_txt]193 194 # phone duration loss195 if hparams['dur_loss'] == 'mse':196 losses['pdur'] = F.mse_loss(dur_pred, (dur_gt + 1).log(), reduction='none')197 losses['pdur'] = (losses['pdur'] * nonpadding).sum() / nonpadding.sum()198 dur_pred = (dur_pred.exp() - 1).clamp(min=0)199 elif hparams['dur_loss'] == 'mog':200 return NotImplementedError201 elif hparams['dur_loss'] == 'crf':202 losses['pdur'] = -self.model.dur_predictor.crf(203 dur_pred, dur_gt.long().clamp(min=0, max=31), mask=nonpadding > 0, reduction='mean')204 losses['pdur'] = losses['pdur'] * hparams['lambda_ph_dur']205 206 # use linear scale for sent and word duration207 if hparams['lambda_word_dur'] > 0:208 word_id = (is_sil.cumsum(-1) * (1 - is_sil)).long()209 word_dur_p = dur_pred.new_zeros([B, word_id.max() + 1]).scatter_add(1, word_id, dur_pred)[:, 1:]210 word_dur_g = dur_gt.new_zeros([B, word_id.max() + 1]).scatter_add(1, word_id, dur_gt)[:, 1:]211 wdur_loss = F.mse_loss((word_dur_p + 1).log(), (word_dur_g + 1).log(), reduction='none')212 word_nonpadding = (word_dur_g > 0).float()213 wdur_loss = (wdur_loss * word_nonpadding).sum() / word_nonpadding.sum()214 losses['wdur'] = wdur_loss * hparams['lambda_word_dur']215 if hparams['lambda_sent_dur'] > 0:216 sent_dur_p = dur_pred.sum(-1)217 sent_dur_g = dur_gt.sum(-1)218 sdur_loss = F.mse_loss((sent_dur_p + 1).log(), (sent_dur_g + 1).log(), reduction='mean')219 losses['sdur'] = sdur_loss.mean() * hparams['lambda_sent_dur']220 221 def add_pitch_loss(self, output, sample, losses):222 if hparams['pitch_type'] == 'ph':223 nonpadding = (sample['txt_tokens'] != 0).float()224 pitch_loss_fn = F.l1_loss if hparams['pitch_loss'] == 'l1' else F.mse_loss225 losses['f0'] = (pitch_loss_fn(output['pitch_pred'][:, :, 0], sample['f0'],226 reduction='none') * nonpadding).sum() \227 / nonpadding.sum() * hparams['lambda_f0']228 return229 mel2ph = sample['mel2ph'] # [B, T_s]230 f0 = sample['f0']231 uv = sample['uv']232 nonpadding = (mel2ph != 0).float()233 if hparams['pitch_type'] == 'cwt':234 cwt_spec = sample[f'cwt_spec']235 f0_mean = sample['f0_mean']236 f0_std = sample['f0_std']237 cwt_pred = output['cwt'][:, :, :10]238 f0_mean_pred = output['f0_mean']239 f0_std_pred = output['f0_std']240 losses['C'] = self.cwt_loss(cwt_pred, cwt_spec) * hparams['lambda_f0']241 if hparams['use_uv']:242 assert output['cwt'].shape[-1] == 11243 uv_pred = output['cwt'][:, :, -1]244 losses['uv'] = (F.binary_cross_entropy_with_logits(uv_pred, uv, reduction='none') * nonpadding) \245 .sum() / nonpadding.sum() * hparams['lambda_uv']246 losses['f0_mean'] = F.l1_loss(f0_mean_pred, f0_mean) * hparams['lambda_f0']247 losses['f0_std'] = F.l1_loss(f0_std_pred, f0_std) * hparams['lambda_f0']248 if hparams['cwt_add_f0_loss']:249 f0_cwt_ = self.model.cwt2f0_norm(cwt_pred, f0_mean_pred, f0_std_pred, mel2ph)250 self.add_f0_loss(f0_cwt_[:, :, None], f0, uv, losses, nonpadding=nonpadding)251 elif hparams['pitch_type'] == 'frame':252 self.add_f0_loss(output['pitch_pred'], f0, uv, losses, nonpadding=nonpadding)253 254 def add_f0_loss(self, p_pred, f0, uv, losses, nonpadding):255 assert p_pred[..., 0].shape == f0.shape256 if hparams['use_uv']:257 assert p_pred[..., 1].shape == uv.shape258 losses['uv'] = (F.binary_cross_entropy_with_logits(259 p_pred[:, :, 1], uv, reduction='none') * nonpadding).sum() \260 / nonpadding.sum() * hparams['lambda_uv']261 nonpadding = nonpadding * (uv == 0).float()262 263 f0_pred = p_pred[:, :, 0]264 if hparams['pitch_loss'] in ['l1', 'l2']:265 pitch_loss_fn = F.l1_loss if hparams['pitch_loss'] == 'l1' else F.mse_loss266 losses['f0'] = (pitch_loss_fn(f0_pred, f0, reduction='none') * nonpadding).sum() \267 / nonpadding.sum() * hparams['lambda_f0']268 elif hparams['pitch_loss'] == 'ssim':269 return NotImplementedError270 271 def cwt_loss(self, cwt_p, cwt_g):272 if hparams['cwt_loss'] == 'l1':273 return F.l1_loss(cwt_p, cwt_g)274 if hparams['cwt_loss'] == 'l2':275 return F.mse_loss(cwt_p, cwt_g)276 if hparams['cwt_loss'] == 'ssim':277 return self.ssim_loss(cwt_p, cwt_g, 20)278 279 def add_energy_loss(self, energy_pred, energy, losses):280 nonpadding = (energy != 0).float()281 loss = (F.mse_loss(energy_pred, energy, reduction='none') * nonpadding).sum() / nonpadding.sum()282 loss = loss * hparams['lambda_energy']283 losses['e'] = loss284 285 286 ############287 # validation plots288 ############289 def plot_mel(self, batch_idx, spec, spec_out, name=None):290 spec_cat = torch.cat([spec, spec_out], -1)291 name = f'mel_{batch_idx}' if name is None else name292 vmin = hparams['mel_vmin']293 vmax = hparams['mel_vmax']294 self.logger.experiment.add_figure(name, spec_to_figure(spec_cat[0], vmin, vmax), self.global_step)295 296 def plot_dur(self, batch_idx, sample, model_out):297 T_txt = sample['txt_tokens'].shape[1]298 dur_gt = mel2ph_to_dur(sample['mel2ph'], T_txt)[0]299 dur_pred = self.model.dur_predictor.out2dur(model_out['dur']).float()300 txt = self.phone_encoder.decode(sample['txt_tokens'][0].cpu().numpy())301 txt = txt.split(" ")302 self.logger.experiment.add_figure(303 f'dur_{batch_idx}', dur_to_figure(dur_gt, dur_pred, txt), self.global_step)304 305 def plot_pitch(self, batch_idx, sample, model_out):306 f0 = sample['f0']307 if hparams['pitch_type'] == 'ph':308 mel2ph = sample['mel2ph']309 f0 = self.expand_f0_ph(f0, mel2ph)310 f0_pred = self.expand_f0_ph(model_out['pitch_pred'][:, :, 0], mel2ph)311 self.logger.experiment.add_figure(312 f'f0_{batch_idx}', f0_to_figure(f0[0], None, f0_pred[0]), self.global_step)313 return314 f0 = denorm_f0(f0, sample['uv'], hparams)315 if hparams['pitch_type'] == 'cwt':316 # cwt317 cwt_out = model_out['cwt']318 cwt_spec = cwt_out[:, :, :10]319 cwt = torch.cat([cwt_spec, sample['cwt_spec']], -1)320 self.logger.experiment.add_figure(f'cwt_{batch_idx}', spec_to_figure(cwt[0]), self.global_step)321 # f0322 f0_pred = cwt2f0(cwt_spec, model_out['f0_mean'], model_out['f0_std'], hparams['cwt_scales'])323 if hparams['use_uv']:324 assert cwt_out.shape[-1] == 11325 uv_pred = cwt_out[:, :, -1] > 0326 f0_pred[uv_pred > 0] = 0327 f0_cwt = denorm_f0(sample['f0_cwt'], sample['uv'], hparams)328 self.logger.experiment.add_figure(329 f'f0_{batch_idx}', f0_to_figure(f0[0], f0_cwt[0], f0_pred[0]), self.global_step)330 elif hparams['pitch_type'] == 'frame':331 # f0332 uv_pred = model_out['pitch_pred'][:, :, 1] > 0333 pitch_pred = denorm_f0(model_out['pitch_pred'][:, :, 0], uv_pred, hparams)334 self.logger.experiment.add_figure(335 f'f0_{batch_idx}', f0_to_figure(f0[0], None, pitch_pred[0]), self.global_step)336 337 ############338 # infer339 ############340 def test_step(self, sample, batch_idx):341 spk_embed = sample.get('spk_embed') if not hparams['use_spk_id'] else sample.get('spk_ids')342 txt_tokens = sample['txt_tokens']343 mel2ph, uv, f0 = None, None, None344 ref_mels = None345 if hparams['profile_infer']:346 pass347 else:348 if hparams['use_gt_dur']:349 mel2ph = sample['mel2ph']350 if hparams['use_gt_f0']:351 f0 = sample['f0']352 uv = sample['uv']353 print('Here using gt f0!!')354 if hparams.get('use_midi') is not None and hparams['use_midi']:355 outputs = self.model(356 txt_tokens, spk_embed=spk_embed, mel2ph=mel2ph, f0=f0, uv=uv, ref_mels=ref_mels, infer=True,357 pitch_midi=sample['pitch_midi'], midi_dur=sample.get('midi_dur'), is_slur=sample.get('is_slur'))358 else:359 outputs = self.model(360 txt_tokens, spk_embed=spk_embed, mel2ph=mel2ph, f0=f0, uv=uv, ref_mels=ref_mels, infer=True)361 sample['outputs'] = self.model.out2mel(outputs['mel_out'])362 sample['mel2ph_pred'] = outputs['mel2ph']363 if hparams.get('pe_enable') is not None and hparams['pe_enable']:364 sample['f0'] = self.pe(sample['mels'])['f0_denorm_pred'] # pe predict from GT mel365 sample['f0_pred'] = self.pe(sample['outputs'])['f0_denorm_pred'] # pe predict from Pred mel366 else:367 sample['f0'] = denorm_f0(sample['f0'], sample['uv'], hparams)368 sample['f0_pred'] = outputs.get('f0_denorm')369 return self.after_infer(sample)370 371 def after_infer(self, predictions):372 if self.saving_result_pool is None and not hparams['profile_infer']:373 self.saving_result_pool = Pool(min(int(os.getenv('N_PROC', os.cpu_count())), 16))374 self.saving_results_futures = []375 predictions = utils.unpack_dict_to_list(predictions)376 t = tqdm(predictions)377 for num_predictions, prediction in enumerate(t):378 for k, v in prediction.items():379 if type(v) is torch.Tensor:380 prediction[k] = v.cpu().numpy()381 382 item_name = prediction.get('item_name')383 text = prediction.get('text').replace(":", "%3A")[:80]384 385 # remove paddings386 mel_gt = prediction["mels"]387 mel_gt_mask = np.abs(mel_gt).sum(-1) > 0388 mel_gt = mel_gt[mel_gt_mask]389 mel2ph_gt = prediction.get("mel2ph")390 mel2ph_gt = mel2ph_gt[mel_gt_mask] if mel2ph_gt is not None else None391 mel_pred = prediction["outputs"]392 mel_pred_mask = np.abs(mel_pred).sum(-1) > 0393 mel_pred = mel_pred[mel_pred_mask]394 mel_gt = np.clip(mel_gt, hparams['mel_vmin'], hparams['mel_vmax'])395 mel_pred = np.clip(mel_pred, hparams['mel_vmin'], hparams['mel_vmax'])396 397 mel2ph_pred = prediction.get("mel2ph_pred")398 if mel2ph_pred is not None:399 if len(mel2ph_pred) > len(mel_pred_mask):400 mel2ph_pred = mel2ph_pred[:len(mel_pred_mask)]401 mel2ph_pred = mel2ph_pred[mel_pred_mask]402 403 f0_gt = prediction.get("f0")404 f0_pred = prediction.get("f0_pred")405 if f0_pred is not None:406 f0_gt = f0_gt[mel_gt_mask]407 if len(f0_pred) > len(mel_pred_mask):408 f0_pred = f0_pred[:len(mel_pred_mask)]409 f0_pred = f0_pred[mel_pred_mask]410 411 str_phs = None412 if self.phone_encoder is not None and 'txt_tokens' in prediction:413 str_phs = self.phone_encoder.decode(prediction['txt_tokens'], strip_padding=True)414 gen_dir = os.path.join(hparams['work_dir'],415 f'generated_{self.trainer.global_step}_{hparams["gen_dir_name"]}')416 wav_pred = self.vocoder.spec2wav(mel_pred, f0=f0_pred)417 if not hparams['profile_infer']:418 os.makedirs(gen_dir, exist_ok=True)419 os.makedirs(f'{gen_dir}/wavs', exist_ok=True)420 os.makedirs(f'{gen_dir}/plot', exist_ok=True)421 os.makedirs(os.path.join(hparams['work_dir'], 'P_mels_npy'), exist_ok=True)422 os.makedirs(os.path.join(hparams['work_dir'], 'G_mels_npy'), exist_ok=True)423 self.saving_results_futures.append(424 self.saving_result_pool.apply_async(self.save_result, args=[425 wav_pred, mel_pred, 'P', item_name, text, gen_dir, str_phs, mel2ph_pred, f0_gt, f0_pred]))426 427 if mel_gt is not None and hparams['save_gt']:428 wav_gt = self.vocoder.spec2wav(mel_gt, f0=f0_gt)429 self.saving_results_futures.append(430 self.saving_result_pool.apply_async(self.save_result, args=[431 wav_gt, mel_gt, 'G', item_name, text, gen_dir, str_phs, mel2ph_gt, f0_gt, f0_pred]))432 if hparams['save_f0']:433 import matplotlib.pyplot as plt434 # f0_pred_, _ = get_pitch(wav_pred, mel_pred, hparams)435 f0_pred_ = f0_pred436 f0_gt_, _ = get_pitch(wav_gt, mel_gt, hparams)437 fig = plt.figure()438 plt.plot(f0_pred_, label=r'$f0_P$')439 plt.plot(f0_gt_, label=r'$f0_G$')440 if hparams.get('pe_enable') is not None and hparams['pe_enable']:441 # f0_midi = prediction.get("f0_midi")442 # f0_midi = f0_midi[mel_gt_mask]443 # plt.plot(f0_midi, label=r'$f0_M$')444 pass445 plt.legend()446 plt.tight_layout()447 plt.savefig(f'{gen_dir}/plot/[F0][{item_name}]{text}.png', format='png')448 plt.close(fig)449 450 t.set_description(451 f"Pred_shape: {mel_pred.shape}, gt_shape: {mel_gt.shape}")452 else:453 if 'gen_wav_time' not in self.stats:454 self.stats['gen_wav_time'] = 0455 self.stats['gen_wav_time'] += len(wav_pred) / hparams['audio_sample_rate']456 print('gen_wav_time: ', self.stats['gen_wav_time'])457 458 return {}459 460 @staticmethod461 def save_result(wav_out, mel, prefix, item_name, text, gen_dir, str_phs=None, mel2ph=None, gt_f0=None, pred_f0=None):462 item_name = item_name.replace('/', '-')463 base_fn = f'[{item_name}][{prefix}]'464 465 if text is not None:466 base_fn += text467 base_fn += ('-' + hparams['exp_name'])468 np.save(os.path.join(hparams['work_dir'], f'{prefix}_mels_npy', item_name), mel)469 audio.save_wav(wav_out, f'{gen_dir}/wavs/{base_fn}.wav', hparams['audio_sample_rate'],470 norm=hparams['out_wav_norm'])471 fig = plt.figure(figsize=(14, 10))472 spec_vmin = hparams['mel_vmin']473 spec_vmax = hparams['mel_vmax']474 heatmap = plt.pcolor(mel.T, vmin=spec_vmin, vmax=spec_vmax)475 fig.colorbar(heatmap)476 if hparams.get('pe_enable') is not None and hparams['pe_enable']:477 gt_f0 = (gt_f0 - 100) / (800 - 100) * 80 * (gt_f0 > 0)478 pred_f0 = (pred_f0 - 100) / (800 - 100) * 80 * (pred_f0 > 0)479 plt.plot(pred_f0, c='white', linewidth=1, alpha=0.6)480 plt.plot(gt_f0, c='red', linewidth=1, alpha=0.6)481 else:482 f0, _ = get_pitch(wav_out, mel, hparams)483 f0 = (f0 - 100) / (800 - 100) * 80 * (f0 > 0)484 plt.plot(f0, c='white', linewidth=1, alpha=0.6)485 if mel2ph is not None and str_phs is not None:486 decoded_txt = str_phs.split(" ")487 dur = mel2ph_to_dur(torch.LongTensor(mel2ph)[None, :], len(decoded_txt))[0].numpy()488 dur = [0] + list(np.cumsum(dur))489 for i in range(len(dur) - 1):490 shift = (i % 20) + 1491 plt.text(dur[i], shift, decoded_txt[i])492 plt.hlines(shift, dur[i], dur[i + 1], colors='b' if decoded_txt[i] != '|' else 'black')493 plt.vlines(dur[i], 0, 5, colors='b' if decoded_txt[i] != '|' else 'black',494 alpha=1, linewidth=1)495 plt.tight_layout()496 plt.savefig(f'{gen_dir}/plot/{base_fn}.png', format='png', dpi=1000)497 plt.close(fig)498 499 ##############500 # utils501 ##############502 @staticmethod503 def expand_f0_ph(f0, mel2ph):504 f0 = denorm_f0(f0, None, hparams)505 f0 = F.pad(f0, [1, 0])506 f0 = torch.gather(f0, 1, mel2ph) # [B, T_mel]507 return f0508 509 510if __name__ == '__main__':511 FastSpeech2Task.start()512 