| 188 | |
| 189 | @torch.inference_mode() |
| 190 | def sample( |
| 191 | self, |
| 192 | batch_size: int, |
| 193 | num_steps: int, |
| 194 | progress: bool = True, |
| 195 | rng: list[torch.Generator] | torch.Generator | None = None, |
| 196 | return_noise: bool = False, |
| 197 | mode: Literal["ddpm", "ddim"] = "ddpm", |
| 198 | text_features: Tensor | None = None, |
| 199 | text_null_features: Tensor | None = None, |
| 200 | ): |
| 201 | noise = self.randn(batch_size, *self.sampling_shape, rng=rng, device=self.device) |
| 202 | x = noise |
| 203 | |
| 204 | tqdm_kwargs = dict(desc="sampling", leave=False, disable=not progress) |
| 205 | for timestep in tqdm(list(reversed(range(num_steps))), **tqdm_kwargs): |
| 206 | timesteps = torch.full((batch_size,), timestep, device=self.device).long() |
| 207 | x = self.p_sample( |
| 208 | x, |
| 209 | timesteps, |
| 210 | text_features=text_features, |
| 211 | text_null_features=text_null_features, |
| 212 | mode=mode |
| 213 | ) |
| 214 | |
| 215 | return noise, x if return_noise else x |