(self)
| 21 | 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 = None |
| 36 | for k, v in self.model.fs2.named_parameters(): |
| 37 | if not 'predictor' in k: |
| 38 | v.requires_grad = False |
| 39 | |
| 40 | def build_optimizer(self, model): |
| 41 | self.optimizer = optimizer = torch.optim.AdamW( |
nothing calls this directly
no test coverage detected