如何显示自编码器生成的重构图像?
如何显示自编码器生成的重构图像?
嗨,你的问题其实很常见,就是图像维度的小细节没处理好!你训练好的自编码器其实已经输出了正确的重构图像,但plt.imshow()不买账的原因是:预测结果多了一个通道维度。
你的predictions[i]形状是(32, 32, 1)(因为训练前给图像加了通道维度适配CNN输入要求),但plt.imshow()默认需要的是2D数组(比如(32,32)),多余的单通道维度会导致显示异常。
下面给你两种解决方案,还附带更实用的对比显示技巧:
方案1:修复单张重构图像的显示
只需要在显示时去掉最后一个通道维度就行,用np.squeeze()压缩冗余维度,或者直接索引[:,:,0]:
for i in range(10): plt.figure(figsize=(4, 4)) # 用squeeze去掉通道维度,或者换成predictions[i][:,:,0] plt.imshow(np.squeeze(predictions[i]), cmap="gray_r") plt.title(f"重构图像 {i+1}") plt.axis('off') # 关闭坐标轴让图像更整洁 plt.show()
方案2:对比显示原始图像和重构图像(更推荐)
这样你能直接直观地看到自编码器的重构效果,比单独看重构图像有用多了:
# 选取前10组测试样本和对应的重构结果 num_samples = 10 original_imgs = x_test[:num_samples] reconstructed_imgs = predictions[:num_samples] # 创建画布,设置合适的尺寸 plt.figure(figsize=(20, 4)) for i in range(num_samples): # 绘制原始图像 ax = plt.subplot(2, num_samples, i + 1) plt.imshow(np.squeeze(original_imgs[i]), cmap="gray_r") plt.title("原始图像") plt.axis('off') # 绘制重构图像 ax = plt.subplot(2, num_samples, i + 1 + num_samples) plt.imshow(np.squeeze(reconstructed_imgs[i]), cmap="gray_r") plt.title("重构图像") plt.axis('off') # 自动调整子图间距 plt.tight_layout() plt.show()
补充说明
你之前的代码里用了figsize=(20,3),这个比例太扁了,换成正方形的画布更适合显示MNIST这类方形图像。另外,关闭坐标轴可以让图像的展示更聚焦,不会被坐标轴分散注意力。
备注:内容来源于stack exchange,提问作者Steven Dascoli
相关产品推荐
相关产品推荐

