(self, batch, N=4, n_row=2, sample=False, ddim_steps=50, ddim_eta=0.0, return_keys=None,
quantize_denoised=True, inpaint=True, plot_denoise_rows=False, plot_progressive_rows=True,
plot_diffusion_rows=False, unconditional_guidance_scale=9.0, unconditional_guidance_label=None,
use_ema_scope=True,
**kwargs)
| 376 | |
| 377 | @torch.no_grad() |
| 378 | def log_images(self, batch, N=4, n_row=2, sample=False, ddim_steps=50, ddim_eta=0.0, return_keys=None, |
| 379 | quantize_denoised=True, inpaint=True, plot_denoise_rows=False, plot_progressive_rows=True, |
| 380 | plot_diffusion_rows=False, unconditional_guidance_scale=9.0, unconditional_guidance_label=None, |
| 381 | use_ema_scope=True, |
| 382 | **kwargs): |
| 383 | use_ddim = ddim_steps is not None |
| 384 | |
| 385 | log = dict() |
| 386 | z, c = self.get_input(batch, self.first_stage_key, bs=N) |
| 387 | c_cat_mask, c_cat_image, c = c["c_concat_mask"][0][:N], c["c_concat_image"][0][:N], c["c_crossattn"][0][:N] |
| 388 | N = min(z.shape[0], N) |
| 389 | n_row = min(z.shape[0], n_row) |
| 390 | log["control_mask"] = c_cat_mask * 2.0 - 1.0 |
| 391 | log["control_image"] = c_cat_image * 2.0 - 1.0 |
| 392 | # log["conditioning"] = log_txt_as_img((384, 384), batch[self.cond_stage_key], size=16) |
| 393 | log["conditioning"] = log_txt_as_img((512, 512), batch[self.cond_stage_key], size=16) |
| 394 | |
| 395 | if plot_diffusion_rows: |
| 396 | # get diffusion row |
| 397 | diffusion_row = list() |
| 398 | z_start = z[:n_row] |
| 399 | for t in range(self.num_timesteps): |
| 400 | if t % self.log_every_t == 0 or t == self.num_timesteps - 1: |
| 401 | t = repeat(torch.tensor([t]), '1 -> b', b=n_row) |
| 402 | t = t.to(self.device).long() |
| 403 | noise = torch.randn_like(z_start) |
| 404 | z_noisy = self.q_sample(x_start=z_start, t=t, noise=noise) |
| 405 | diffusion_row.append(self.decode_first_stage(z_noisy)) |
| 406 | |
| 407 | diffusion_row = torch.stack(diffusion_row) # n_log_step, n_row, C, H, W |
| 408 | diffusion_grid = rearrange(diffusion_row, 'n b c h w -> b n c h w') |
| 409 | diffusion_grid = rearrange(diffusion_grid, 'b n c h w -> (b n) c h w') |
| 410 | diffusion_grid = make_grid(diffusion_grid, nrow=diffusion_row.shape[0]) |
| 411 | log["diffusion_row"] = diffusion_grid |
| 412 | |
| 413 | if sample: |
| 414 | # get denoise row |
| 415 | samples, z_denoise_row = self.sample_log(cond={"c_concat": [c_cat_mask], "c_crossattn": [c]}, |
| 416 | batch_size=N, ddim=use_ddim, |
| 417 | ddim_steps=ddim_steps, eta=ddim_eta) |
| 418 | x_samples = self.decode_first_stage(samples) |
| 419 | log["samples"] = x_samples |
| 420 | if plot_denoise_rows: |
| 421 | denoise_grid = self._get_denoise_row_from_list(z_denoise_row) |
| 422 | log["denoise_row"] = denoise_grid |
| 423 | |
| 424 | if unconditional_guidance_scale > 1.0: |
| 425 | uc_cross = self.get_unconditional_conditioning(N) |
| 426 | uc_cat = c_cat_mask # torch.zeros_like(c_cat) |
| 427 | uc_full = {"c_concat": [uc_cat], "c_crossattn": [uc_cross]} |
| 428 | samples_cfg, _ = self.sample_log(cond={"c_concat": [c_cat_mask], "c_crossattn": [c]}, |
| 429 | batch_size=N, ddim=use_ddim, |
| 430 | ddim_steps=ddim_steps, eta=ddim_eta, |
| 431 | unconditional_guidance_scale=unconditional_guidance_scale, |
| 432 | unconditional_conditioning=uc_full, |
| 433 | ) |
| 434 | x_samples_cfg = self.decode_first_stage(samples_cfg) |
| 435 | log[f"samples_cfg_scale_{unconditional_guidance_scale:.2f}_mask"] = x_samples_cfg |
no test coverage detected