| 274 | ) |
| 275 | |
| 276 | def p_losses(self, x_start, t, cond, noise=None, nonpadding=None): |
| 277 | noise = default(noise, lambda: torch.randn_like(x_start)) |
| 278 | |
| 279 | x_noisy = self.q_sample(x_start=x_start, t=t, noise=noise) |
| 280 | x_recon = self.denoise_fn(x_noisy, t, cond) |
| 281 | |
| 282 | if self.loss_type == 'l1': |
| 283 | if nonpadding is not None: |
| 284 | loss = ((noise - x_recon).abs() * nonpadding.unsqueeze(1)).mean() |
| 285 | else: |
| 286 | # print('are you sure w/o nonpadding?') |
| 287 | loss = (noise - x_recon).abs().mean() |
| 288 | |
| 289 | elif self.loss_type == 'l2': |
| 290 | loss = F.mse_loss(noise, x_recon) |
| 291 | else: |
| 292 | raise NotImplementedError() |
| 293 | |
| 294 | return loss |
| 295 | |
| 296 | def forward(self, txt_tokens, mel2ph=None, spk_embed=None, |
| 297 | ref_mels=None, f0=None, uv=None, energy=None, infer=False): |