Sample from the model
(self, denoise_fn, data, t, noise_fn, clip_denoised=False, return_pred_xstart=False, use_var=True)
| 185 | ''' samples ''' |
| 186 | |
| 187 | def p_sample(self, denoise_fn, data, t, noise_fn, clip_denoised=False, return_pred_xstart=False, use_var=True): |
| 188 | """ |
| 189 | Sample from the model |
| 190 | """ |
| 191 | model_mean, _, model_log_variance, pred_xstart = self.p_mean_variance(denoise_fn, data=data, t=t, clip_denoised=clip_denoised, |
| 192 | return_pred_xstart=True) |
| 193 | noise = noise_fn(size=data.shape, dtype=data.dtype, device=data.device) |
| 194 | assert noise.shape == data.shape |
| 195 | # no noise when t == 0 |
| 196 | nonzero_mask = torch.reshape(1 - (t == 0).float(), [data.shape[0]] + [1] * (len(data.shape) - 1)) |
| 197 | |
| 198 | sample = model_mean |
| 199 | if use_var: |
| 200 | sample = sample + nonzero_mask * torch.exp(0.5 * model_log_variance) * noise |
| 201 | assert sample.shape == pred_xstart.shape |
| 202 | return (sample, pred_xstart) if return_pred_xstart else sample |
| 203 | |
| 204 | |
| 205 | def p_sample_loop(self, denoise_fn, shape, device, |
no test coverage detected