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

Function vis_density_GMM

utils.py:149–172  ·  view source on GitHub ↗
(model, config)

Source from the content-addressed store, hash-verified

147
148
149def vis_density_GMM(model, config):
150 fig, ax = plt.subplots()
151 x_density, y_density = 500, 500
152 xses = np.linspace(-10, 10, x_density)
153 yses = np.linspace(-10, 10, y_density)
154 xy = torch.tensor([[[x, y] for x in xses]
155 for y in yses]).view(-1, 2).cuda().float()
156 log_density_values = model.log_prob(xy)
157 log_density_values = log_density_values.detach().view(
158 x_density, y_density).cpu().numpy()
159 dx = (xses[1] - xses[0]) / 2
160 dy = (yses[1] - yses[0]) / 2
161 extent = [xses[0] - dx, xses[-1] + dx, yses[0] - dy, yses[-1] + dy]
162 im = ax.imshow(np.exp(log_density_values),
163 extent=extent,
164 origin='lower',
165 cmap='viridis')
166 divider = make_axes_locatable(ax)
167 cax = divider.append_axes('right', size='5%', pad=0.05)
168 cb = fig.colorbar(im, cax=cax)
169 cb.set_label('probability density')
170 plt.show()
171 plt.savefig(f"{config['exp_folder']}/GMM_density.png", bbox_inches='tight')
172 plt.close()
173
174
175def vis_density_GRBM(model, config, epoch=None):

Callers 1

create_datasetFunction · 0.90

Calls 1

log_probMethod · 0.80

Tested by

no test coverage detected