(model, epoch, config, tag=None, is_show_gif=True)
| 96 | |
| 97 | |
| 98 | def visualize_sampling(model, epoch, config, tag=None, is_show_gif=True): |
| 99 | tag = '' if tag is None else tag |
| 100 | B, C, H, W = config['sampling_batch_size'], config['channel'], config[ |
| 101 | 'height'], config['width'] |
| 102 | v_init = torch.randn(B, C, H, W).cuda() |
| 103 | v_list = model.sampling(v_init, |
| 104 | num_steps=config['sampling_steps'], |
| 105 | save_gap=config['sampling_gap']) |
| 106 | |
| 107 | if 'GMM' in config['dataset']: |
| 108 | samples = v_list[-1][1].view(B, -1).cpu().numpy() |
| 109 | vis_2D_samples(samples, config, tags=f'{epoch:05d}') |
| 110 | vis_density_GRBM(model, config, epoch=epoch) |
| 111 | else: |
| 112 | if is_show_gif: |
| 113 | v_list = unnormalize_img_tuple(v_list, config['img_mean'], |
| 114 | config['img_std']) |
| 115 | save_gif_fancy( |
| 116 | v_list, config['sampling_nrow'], |
| 117 | f"{config['exp_folder']}/sample_imgs_epoch_{epoch:05d}{tag}.gif") |
| 118 | img_vis = v_list[-1][1] |
| 119 | else: |
| 120 | if isinstance(config['img_std'], torch.Tensor): |
| 121 | mean = config['img_mean'].view(1, -1, 1, 1).cuda() |
| 122 | std = config['img_std'].view(1, -1, 1, 1).cuda() |
| 123 | else: |
| 124 | mean = config['img_mean'] |
| 125 | std = config['img_std'] |
| 126 | |
| 127 | img_vis = (v_list[-1][1] * std + mean).clamp(min=0, max=1) |
| 128 | |
| 129 | utils.save_image( |
| 130 | utils.make_grid(img_vis, |
| 131 | nrow=config['sampling_nrow'], |
| 132 | normalize=False, |
| 133 | padding=1, |
| 134 | pad_value=1.0).cpu(), |
| 135 | f"{config['exp_folder']}/sample_imgs_epoch_{epoch:05d}{tag}.png") |
| 136 | |
| 137 | |
| 138 | def vis_2D_samples(samples, config, tags=None): |
no test coverage detected