You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.23 17:15:03