CoolFace
Apppublic

Silentlin/DiffSinger

sourceHugging Faceupdated 3y agoView on Hugging Face
89likes
diffspeech_task.py123 linesDownload Raw Back to usr
1import torch2 3import utils4from utils.hparams import hparams5from .diff.net import DiffNet6from .diff.shallow_diffusion_tts import GaussianDiffusion7from .task import DiffFsTask8from vocoders.base_vocoder import get_vocoder_cls, BaseVocoder9from utils.pitch_utils import denorm_f010from tasks.tts.fs2_utils import FastSpeechDataset11 12DIFF_DECODERS = {13    'wavenet': lambda hp: DiffNet(hp['audio_num_mel_bins']),14}15 16 17class DiffSpeechTask(DiffFsTask):18    def __init__(self):19        super(DiffSpeechTask, self).__init__()20        self.dataset_cls = FastSpeechDataset21        self.vocoder: BaseVocoder = get_vocoder_cls(hparams)()22 23    def build_tts_model(self):24        mel_bins = hparams['audio_num_mel_bins']25        self.model = GaussianDiffusion(26            phone_encoder=self.phone_encoder,27            out_dims=mel_bins, denoise_fn=DIFF_DECODERS[hparams['diff_decoder_type']](hparams),28            timesteps=hparams['timesteps'],29            K_step=hparams['K_step'],30            loss_type=hparams['diff_loss_type'],31            spec_min=hparams['spec_min'], spec_max=hparams['spec_max'],32        )33        if hparams['fs2_ckpt'] != '':34            utils.load_ckpt(self.model.fs2, hparams['fs2_ckpt'], 'model', strict=True)35        # self.model.fs2.decoder = None36        for k, v in self.model.fs2.named_parameters():37            if not 'predictor' in k:38                v.requires_grad = False39 40    def build_optimizer(self, model):41        self.optimizer = optimizer = torch.optim.AdamW(42            filter(lambda p: p.requires_grad, model.parameters()),43            lr=hparams['lr'],44            betas=(hparams['optimizer_adam_beta1'], hparams['optimizer_adam_beta2']),45            weight_decay=hparams['weight_decay'])46        return optimizer47 48    def run_model(self, model, sample, return_output=False, infer=False):49        txt_tokens = sample['txt_tokens']  # [B, T_t]50        target = sample['mels']  # [B, T_s, 80]51        # mel2ph = sample['mel2ph'] if hparams['use_gt_dur'] else None # [B, T_s]52        mel2ph = sample['mel2ph']53        f0 = sample['f0']54        uv = sample['uv']55        energy = sample['energy']56        # fs2_mel = sample['fs2_mels']57        spk_embed = sample.get('spk_embed') if not hparams['use_spk_id'] else sample.get('spk_ids')58        if hparams['pitch_type'] == 'cwt':59            cwt_spec = sample[f'cwt_spec']60            f0_mean = sample['f0_mean']61            f0_std = sample['f0_std']62            sample['f0_cwt'] = f0 = model.cwt2f0_norm(cwt_spec, f0_mean, f0_std, mel2ph)63 64        output = model(txt_tokens, mel2ph=mel2ph, spk_embed=spk_embed,65                       ref_mels=target, f0=f0, uv=uv, energy=energy, infer=infer)66 67        losses = {}68        if 'diff_loss' in output:69            losses['mel'] = output['diff_loss']70        self.add_dur_loss(output['dur'], mel2ph, txt_tokens, losses=losses)71        if hparams['use_pitch_embed']:72            self.add_pitch_loss(output, sample, losses)73        if hparams['use_energy_embed']:74            self.add_energy_loss(output['energy_pred'], energy, losses)75        if not return_output:76            return losses77        else:78            return losses, output79 80    def validation_step(self, sample, batch_idx):81        outputs = {}82        txt_tokens = sample['txt_tokens']  # [B, T_t]83 84        energy = sample['energy']85        spk_embed = sample.get('spk_embed') if not hparams['use_spk_id'] else sample.get('spk_ids')86        mel2ph = sample['mel2ph']87        f0 = sample['f0']88        uv = sample['uv']89 90        outputs['losses'] = {}91 92        outputs['losses'], model_out = self.run_model(self.model, sample, return_output=True, infer=False)93 94 95        outputs['total_loss'] = sum(outputs['losses'].values())96        outputs['nsamples'] = sample['nsamples']97        outputs = utils.tensors_to_scalars(outputs)98        if batch_idx < hparams['num_valid_plots']:99            # model_out = self.model(100            #     txt_tokens, spk_embed=spk_embed, mel2ph=None, f0=None, uv=None, energy=None, ref_mels=None, infer=True)101            # self.plot_mel(batch_idx, model_out['mel_out'], model_out['fs2_mel'], name=f'diffspeech_vs_fs2_{batch_idx}')102            model_out = self.model(103                txt_tokens, spk_embed=spk_embed, mel2ph=mel2ph, f0=f0, uv=uv, energy=energy, ref_mels=None, infer=True)104            gt_f0 = denorm_f0(sample['f0'], sample['uv'], hparams)105            self.plot_wav(batch_idx, sample['mels'], model_out['mel_out'], is_mel=True, gt_f0=gt_f0, f0=model_out.get('f0_denorm'))106            self.plot_mel(batch_idx, sample['mels'], model_out['mel_out'])107        return outputs108 109    ############110    # validation plots111    ############112    def plot_wav(self, batch_idx, gt_wav, wav_out, is_mel=False, gt_f0=None, f0=None, name=None):113        gt_wav = gt_wav[0].cpu().numpy()114        wav_out = wav_out[0].cpu().numpy()115        gt_f0 = gt_f0[0].cpu().numpy()116        f0 = f0[0].cpu().numpy()117        if is_mel:118            gt_wav = self.vocoder.spec2wav(gt_wav, f0=gt_f0)119            wav_out = self.vocoder.spec2wav(wav_out, f0=f0)120        self.logger.experiment.add_audio(f'gt_{batch_idx}', gt_wav, sample_rate=hparams['audio_sample_rate'], global_step=self.global_step)121        self.logger.experiment.add_audio(f'wav_{batch_idx}', wav_out, sample_rate=hparams['audio_sample_rate'], global_step=self.global_step)122 123