MCPcopy Create free account
hub / github.com/VisionXLab/OF-Diff / log_images

Method log_images

cldm/cldm.py:378–448  ·  view source on GitHub ↗
(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)

Source from the content-addressed store, hash-verified

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

Callers 1

log_imgMethod · 0.45

Calls 7

get_inputMethod · 0.95
sample_logMethod · 0.95
log_txt_as_imgFunction · 0.90
decode_first_stageMethod · 0.80
q_sampleMethod · 0.45

Tested by

no test coverage detected