(
self,
model,
x_latent,
t_start=None,
ddim_steps=100,
use_original_steps=True,
model_kwargs=None,
progress=True,
clip_denoised=False
)
| 282 | |
| 283 | @torch.no_grad() |
| 284 | def decode( |
| 285 | self, |
| 286 | model, |
| 287 | x_latent, |
| 288 | t_start=None, |
| 289 | ddim_steps=100, |
| 290 | use_original_steps=True, |
| 291 | model_kwargs=None, |
| 292 | progress=True, |
| 293 | clip_denoised=False |
| 294 | ): |
| 295 | bs, dev = x_latent.shape[0], x_latent.device |
| 296 | model_kwargs = model_kwargs or {} |
| 297 | |
| 298 | self.make_schedule(ddim_num_steps=ddim_steps, device=dev, ddim_eta=0.0, verbose=False) |
| 299 | |
| 300 | timesteps = np.arange(self.ddpm_num_timesteps) if use_original_steps else self.ddim_timesteps |
| 301 | if t_start is not None: |
| 302 | timesteps = timesteps[:t_start] |
| 303 | |
| 304 | time_range = np.flip(timesteps) |
| 305 | total_steps = timesteps.shape[0] |
| 306 | |
| 307 | iterator = tqdm(time_range, desc='Decoding image', total=total_steps, disable=not progress) |
| 308 | |
| 309 | x_dec = x_latent |
| 310 | for i, step in enumerate(iterator): |
| 311 | index = total_steps - i - 1 |
| 312 | ts = torch.full((bs,), step, device=dev, dtype=torch.long) |
| 313 | |
| 314 | x_dec, _ = self.p_sample_ddim( |
| 315 | model=model, |
| 316 | x=x_dec, |
| 317 | t=ts, |
| 318 | index=index, |
| 319 | model_kwargs=model_kwargs, |
| 320 | use_original_steps=use_original_steps, |
| 321 | clip_denoised=clip_denoised |
| 322 | ) |
| 323 | |
| 324 | return x_dec |
nothing calls this directly
no test coverage detected