(self, x, c, t, clip_denoised=False, repeat_noise=False, return_x0=False, \
temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None, **kwargs)
| 860 | |
| 861 | @torch.no_grad() |
| 862 | def p_sample(self, x, c, t, clip_denoised=False, repeat_noise=False, return_x0=False, \ |
| 863 | temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None, **kwargs): |
| 864 | b, *_, device = *x.shape, x.device |
| 865 | outputs = self.p_mean_variance(x=x, c=c, t=t, clip_denoised=clip_denoised, return_x0=return_x0, \ |
| 866 | score_corrector=score_corrector, corrector_kwargs=corrector_kwargs, **kwargs) |
| 867 | if return_x0: |
| 868 | model_mean, _, model_log_variance, x0 = outputs |
| 869 | else: |
| 870 | model_mean, _, model_log_variance = outputs |
| 871 | |
| 872 | noise = noise_like(x.shape, device, repeat_noise) * temperature |
| 873 | if noise_dropout > 0.: |
| 874 | noise = torch.nn.functional.dropout(noise, p=noise_dropout) |
| 875 | # no noise when t == 0 |
| 876 | nonzero_mask = (1 - (t == 0).float()).reshape(b, *((1,) * (len(x.shape) - 1))) |
| 877 | |
| 878 | if return_x0: |
| 879 | return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise, x0 |
| 880 | else: |
| 881 | return model_mean + nonzero_mask * (0.5 * model_log_variance).exp() * noise |
| 882 | |
| 883 | @torch.no_grad() |
| 884 | def p_sample_loop(self, cond, shape, return_intermediates=False, x_T=None, verbose=True, callback=None, \ |
no test coverage detected