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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 08:10:47