def visualize_class_indices(class_indices_tensor, num_classes, cmap='tab20'):
“”class_indices_tensor的shape是2维的,如[256,256],num_classes为多少根据自己的类别数确定“”
colors = plt.get_cmap(cmap)(np.linspace(0, 1, num_classes))[:, :3]
class_colors = mcolors.ListedColormap(colors)
class_indices_tensor = class_indices_tensor.cpu().numpy() # Convert to NumPy array and move to CPU
#假设有15个类别 下面是GID数据图例(一个例子)
#class_labels = ['工业用地', '城市住宅','农村住宅', '交通用地','稻田', '灌溉地','旱地', '园地','乔木林', '灌木林','自然草地','人工草地', '河流','湖泊', '池塘']
# 创建图像
fig, ax = plt.subplots(figsize=(18, 12))
ax.imshow(class_indices_tensor, cmap=class_colors, vmin=0, vmax=num_classes - 1)
# 创建Patch对象列表
legend_patches = [mpatches.Patch(color=class_colors(i), label=class_labels[i]) for i in range(num_classes)]
# 添加legend
ax.legend(handles=legend_patches, bbox_to_anchor=(1.05, 1), loc='upper left', borderaxespad=0.,fontsize=30)
ax.set_title('Output', fontsize=30)
# ax.axis('off') # 隐藏坐标轴刻度和标签
return fig
注意点:
1 输入数据尺寸需要是2维的
2 数据中的类别数可以根据自己需要变化
3 多次调用该函数,如果比如同样都是15个类别,那么每次类别对应的颜色是不会变的,如果想要变化,可以更改cmap(函数参数).