(model, config)
| 147 | |
| 148 | |
| 149 | def 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 | |
| 175 | def vis_density_GRBM(model, config, epoch=None): |
no test coverage detected