| 374 | |
| 375 | @torch.no_grad() |
| 376 | def p_sample_loop(self, shape, return_intermediates=False): |
| 377 | device = self.betas.device |
| 378 | b = shape[0] |
| 379 | img = torch.randn(shape, device=device) |
| 380 | intermediates = [img] |
| 381 | for i in tqdm(reversed(range(0, self.num_timesteps)), desc='Sampling t', total=self.num_timesteps): |
| 382 | img = self.p_sample(img, torch.full((b,), i, device=device, dtype=torch.long), |
| 383 | clip_denoised=self.clip_denoised) |
| 384 | if i % self.log_every_t == 0 or i == self.num_timesteps - 1: |
| 385 | intermediates.append(img) |
| 386 | if return_intermediates: |
| 387 | return img, intermediates |
| 388 | return img |
| 389 | |
| 390 | @torch.no_grad() |
| 391 | def sample(self, batch_size=16, return_intermediates=False): |