CoolFace
Apppublic

justyoung/DiffSinger

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
fs2.py512 linesDownload Raw Back to tts
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