如何在Pyplot中显示U-Net输出的MNIST独热编码分割数据
解决plt.imshow显示独热编码分割掩码的问题
你遇到的TypeError: Invalid dimensions for image data是因为plt.imshow只支持2D灰度图(形状为(H,W))或者3D彩色图(形状为(H,W,3)或(H,W,4)),而你的独热编码数组是(H,W,11)的11通道格式,完全不符合它的输入要求。
要实现你想要的MNIST分割可视化效果,我们需要把独热编码转换成单通道的类别索引掩码,再进行可视化。具体步骤如下:
1. 将独热编码转换为类别索引
独热编码中每个像素位置只有一个值为1的通道,对应该像素的类别(0-9是数字,10是空白)。我们可以用np.argmax()提取每个像素对应的类别索引,把(H,W,11)的数组转换成(H,W)的2D数组:
# 提取你的独热编码数组中的第一个样本(可替换为任意样本) sample_mask = one_hot_coded_arr[0] # 形状为(128,128,11),对应你代码生成的数据 segmentation_index = np.argmax(sample_mask, axis=-1) # 转换后形状为(128,128),每个值是0-10
2. 自定义颜色映射(可选但推荐)
为了让不同数字和空白区域有明显区分,我们可以自定义颜色映射:给0-9分配不同的亮色,空白(10)分配黑色:
import matplotlib.colors as mcolors # 用tab10配色(10种颜色对应数字0-9),再添加黑色对应空白区域 color_list = plt.cm.tab10.colors + [(0, 0, 0)] custom_cmap = mcolors.ListedColormap(color_list)
3. 可视化分割掩码
现在用plt.imshow处理转换后的2D索引数组,配合自定义颜色映射即可:
plt.figure(figsize=(6,6)) plt.imshow(segmentation_index, cmap=custom_cmap, vmin=0, vmax=10, interpolation='nearest') plt.axis('off') # 关闭坐标轴 # 可选:添加颜色条标注每个颜色对应的类别 cbar = plt.colorbar(ticks=np.arange(11)) cbar.set_ticklabels(['0', '1', '2', '3', '4', '5', '6', '7', '8', '9', '空白']) plt.show()
整合到你代码中的完整示例
把上述逻辑替换你代码中原有的plt.imshow部分即可:
# 原代码生成one_hot_coded_arr之后的部分 print(one_hot_coded_arr.shape) # 转换独热编码为类别索引 sample_index = np.argmax(one_hot_coded_arr[0], axis=-1) # 自定义颜色映射 import matplotlib.colors as mcolors color_list = plt.cm.tab10.colors + [(0,0,0)] custom_cmap = mcolors.ListedColormap(color_list) # 可视化 plt.imshow(sample_index, cmap=custom_cmap, vmin=0, vmax=10, interpolation='nearest') plt.axis("off") cbar = plt.colorbar(ticks=np.arange(11)) cbar.set_ticklabels(['0','1','2','3','4','5','6','7','8','9','空白']) plt.show()
这样就能生成类似你想要的分割可视化图,每个数字类别用不同颜色区分,空白区域为黑色。
内容的提问来源于stack exchange,提问作者yupthatsme
相关产品推荐
相关产品推荐

