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

如何调整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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 12:42:15