MCPcopy Create free account
hub / github.com/DSL-Lab/GRBM / visualize_sampling

Function visualize_sampling

utils.py:98–135  ·  view source on GitHub ↗
(model, epoch, config, tag=None, is_show_gif=True)

Source from the content-addressed store, hash-verified

96
97
98def 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
138def vis_2D_samples(samples, config, tags=None):

Callers 1

train_modelFunction · 0.90

Calls 5

vis_2D_samplesFunction · 0.85
vis_density_GRBMFunction · 0.85
unnormalize_img_tupleFunction · 0.85
save_gif_fancyFunction · 0.85
samplingMethod · 0.45

Tested by

no test coverage detected