如何调整matplotlib混淆矩阵尺寸以完整显示100类标签并查看误判
100类混淆矩阵显示优化方案
核心问题
原代码的混淆矩阵绘制逻辑是适配10类左右小分类数场景写的,直接套用到100类场景会出现三个硬伤:
- 画布尺寸固定为44英寸,没有足够空间排布100个轴标签和100100的矩阵单元格
- 默认给所有单元格标注数值,10000个数字叠加在一起完全无法辨认
- 轴标签字号、旋转角度、边距没有做多分类适配,标签互相重叠、被画布截断
具体修改步骤
第一步:替换原有的混淆矩阵绘制函数
把代码里的plot_confusion_matrix函数替换成下面适配100类场景的版本:
def plot_confusion_matrix(y_true, y_pred, classes, normalize=False, title=None, cmap=plt.cm.Blues): if not title: if normalize: title = 'Normalized confusion matrix' else: title = 'Confusion matrix, without normalization' # 计算混淆矩阵 cm = confusion_matrix(y_true, y_pred) if normalize: cm = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis] print("Normalized confusion matrix") else: print('Confusion matrix, without normalization') # 调大画布尺寸,100类建议设为20-25英寸正方形,dpi设为100保证基础清晰度 fig, ax = plt.subplots(figsize=(22, 22), dpi=100) im = ax.imshow(cm, interpolation='nearest', cmap=cmap) ax.figure.colorbar(im, ax=ax, fraction=0.046, pad=0.04) # 设置轴标签 ax.set(xticks=np.arange(cm.shape[1]), yticks=np.arange(cm.shape[0]), xticklabels=classes, yticklabels=classes, title=title, ylabel='True label', xlabel='Predicted label') # 调小刻度标签字号,x轴标签旋转90度避免重叠 ax.tick_params(axis='both', which='major', labelsize=7) plt.setp(ax.get_xticklabels(), rotation=90, ha="right", rotation_mode="anchor") # 100类场景下注释掉全量单元格数值标注,避免文字糊成一团,靠颜色深浅即可判断数值大小 # fmt = '.2f' if normalize else 'd' # thresh = cm.max() / 2. # for i in range(cm.shape[0]): # for j in range(cm.shape[1]): # ax.text(j, i, format(cm[i, j], fmt), # ha="center", color="white" # if cm[i, j] > thresh else "black") # 增加边距预留,避免标签被截断 plt.tight_layout(pad=2.0) return ax
第二步:添加高清图保存逻辑(可选)
在调用plot_confusion_matrix之后、plt.show()之前加一行保存代码,导出300dpi的高清图,本地放大后可以清晰看到每个单元格的颜色和对应标签:
# 绘制非归一化混淆矩阵 plot_confusion_matrix(y_true, y_pred, classes=class_names, title='AlexNet Confusion matrix, without normalization') # 保存高清图 plt.savefig('confusion_matrix_raw.png', dpi=300, bbox_inches='tight') plt.show() # 如果需要看比例更均匀的误判分布,可以打开归一化矩阵绘制 plot_confusion_matrix(y_true, y_pred, classes=class_names, normalize=True, title='Normalized confusion matrix') plt.savefig('confusion_matrix_norm.png', dpi=300, bbox_inches='tight') plt.show()
第三步:精准定位误判样本(可选)
如果需要看具体误判的数值,不要依赖图上的文字标注,直接把混淆矩阵导出为csv文件查询即可:
import pandas as pd # 导出非归一化矩阵 pd.DataFrame(confusion_mtx, index=class_names, columns=class_names).to_csv('confusion_matrix_raw.csv') # 导出归一化矩阵 cm_norm = confusion_mtx.astype('float') / confusion_mtx.sum(axis=1)[:, np.newaxis] pd.DataFrame(cm_norm, index=class_names, columns=class_names).to_csv('confusion_matrix_norm.csv')
配合热力图的颜色定位到误判率高的区域后,直接在csv里查对应行列的具体数值即可,效率比在图上找数字高很多。
参数调整说明
- 如果你的类别名称长度普遍超过10个字符,可以把
figsize的数值从22调到25甚至更大 - 如果觉得标签还是挤,可以把
labelsize=7再调小到6,或者每隔2个刻度显示一个标签 - 归一化后的混淆矩阵颜色对比更均匀,更适合观察整体误判分布,非归一化矩阵适合看每个类误判的绝对样本数
内容的提问来源于stack exchange,提问作者Josh
相关产品推荐
相关产品推荐

