(self, denoise_fn, data_start, data_t, t, clip_denoised: bool, return_pred_xstart: bool)
| 294 | '''losses''' |
| 295 | |
| 296 | def _vb_terms_bpd(self, denoise_fn, data_start, data_t, t, clip_denoised: bool, return_pred_xstart: bool): |
| 297 | true_mean, _, true_log_variance_clipped = self.q_posterior_mean_variance(x_start=data_start, x_t=data_t, t=t) |
| 298 | model_mean, _, model_log_variance, pred_xstart = self.p_mean_variance( |
| 299 | denoise_fn, data=data_t, t=t, clip_denoised=clip_denoised, return_pred_xstart=True) |
| 300 | kl = normal_kl(true_mean, true_log_variance_clipped, model_mean, model_log_variance) |
| 301 | kl = kl.mean(dim=list(range(1, len(data_start.shape)))) / np.log(2.) |
| 302 | |
| 303 | return (kl, pred_xstart) if return_pred_xstart else kl |
| 304 | |
| 305 | def p_losses(self, denoise_fn, data_start, t, noise=None): |
| 306 | """ |
no test coverage detected