VAE解码器z_sample维度不匹配报错,如何正确绘制模型输出结果
问题原因
- 从提供的编码器结构可确认,该VAE的隐空间维度为10:全连接层
dense_6输出20维特征后被拆分为两个10维张量,分别对应隐分布的均值和对数方差,重参数化后生成的隐变量z形状为(None, 10),因此解码器要求输入的最后一维必须为10。 - 当前使用的是2维隐空间VAE的绘图代码,生成的
z_sample形状为(1,2),和解码器要求的输入维度不匹配,因此触发报错。
解决方案
10维隐空间无法直接在2D平面完整可视化,通用做法是固定隐空间剩余8个维度为默认值(通常设为0,对应隐分布的均值位置),仅遍历前两个维度生成网格,即可绘制前两维的隐空间变化效果。
仅需修改循环内的z_sample生成逻辑即可:
z_sample = np.array([[xi, yi] + [0]*8])
修改后的完整绘图函数如下:
def plot_latent_space(n=30, figsize=15): digit_size = 28 scale = 1.5 figure = np.zeros((digit_size * n, digit_size * n)) grid_x = np.linspace(-scale, scale, n) grid_y = np.linspace(-scale, scale, n)[::-1] for i, yi in enumerate(grid_y): for j, xi in enumerate(grid_x): # 补全8个0,将z_sample维度扩展为10 z_sample = np.array([[xi, yi] + [0]*8]) x_decoded = vae_decoder(z_sample) digit = tf.reshape(x_decoded[0], shape=(digit_size, digit_size)) figure[ i * digit_size : (i + 1) * digit_size, j * digit_size : (j + 1) * digit_size, ] = digit plt.figure(figsize=(figsize, figsize)) start_range = digit_size // 2 end_range = n * digit_size + start_range pixel_range = np.arange(start_range, end_range, digit_size) sample_range_x = np.round(grid_x, 1) sample_range_y = np.round(grid_y, 1) plt.xticks(pixel_range, sample_range_x) plt.yticks(pixel_range, sample_range_y) plt.xlabel("z[0]") plt.ylabel("z[1]") plt.imshow(figure, cmap="Greys_r") plt.show()
如果需要查看其他维度对的隐空间效果,只需调整z_sample中遍历值的位置,将非遍历维度设为固定值即可。
内容的提问来源于stack exchange,提问作者RustX
相关产品推荐
相关产品推荐

