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

使用Seaborn绘制Sklearn混淆矩阵时的索引错位问题

混淆矩阵热力图显示错误类别问题

测试集真实标签仅包含类别1和3,但使用Seaborn绘制混淆矩阵热力图时,却出现了类别0和2,图表整体下移一行,问题根源在于类别索引不匹配。

原代码及输出

原代码

from sklearn.metrics import confusion_matrix
from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay
import seaborn as sns
import numpy as np
import matplotlib.pyplot as plt
from collections import Counter

cf_matrix = confusion_matrix(y_true, y_pred)
print(Counter(y_pred))
print(Counter(y_true))

cmn = cf_matrix.astype('float') / cf_matrix.sum(axis=1)[:, np.newaxis]
plt.figure(figsize = (15,15))
sns.heatmap(cmn, annot=True, fmt='.1f')

输出

Counter({3: 100489, 12: 11306, 11: 4314, 4: 3303, 8: 2510, 7: 1850, 5: 185, 10: 132, 2: 69})
Counter({3.0: 117955, 1.0: 6203})

问题原因

confusion_matrix默认会根据**所有出现过的类别(包括预测结果中的类别)**生成从0开始的连续整数索引,但你的真实标签只有1和3,预测结果却包含2、4、5等多个类别,导致混淆矩阵的行/列索引与实际类别错位,进而出现无意义的类别0和偏移问题。

解决方法

通过指定confusion_matrix的labels参数,明确混淆矩阵要包含的类别,并在绘制热力图时手动设置刻度标签:

修改后的代码

from sklearn.metrics import confusion_matrix
import seaborn as sns
import numpy as np
import matplotlib.pyplot as plt
from collections import Counter

# 收集真实标签和预测标签中所有出现过的类别,去重并排序
all_classes = sorted(list(set(y_true).union(set(y_pred))))

# 指定labels参数生成与实际类别匹配的混淆矩阵
cf_matrix = confusion_matrix(y_true, y_pred, labels=all_classes)
print(Counter(y_pred))
print(Counter(y_true))

# 计算归一化混淆矩阵
cmn = cf_matrix.astype('float') / cf_matrix.sum(axis=1)[:, np.newaxis]

plt.figure(figsize=(15,15))
# 设置热力图的刻度标签为实际类别,避免索引错位
sns.heatmap(cmn, annot=True, fmt='.1f', xticklabels=all_classes, yticklabels=all_classes)
plt.xlabel('预测类别')
plt.ylabel('真实类别')
plt.show()

可选调整

如果只需要关注测试集真实存在的类别(1和3),可以把all_classes替换为:

all_classes = sorted(list(set(y_true)))

这样混淆矩阵只会包含真实标签中的类别,过滤掉预测结果中出现的其他无关类别。

内容的提问来源于stack exchange,提问作者Leo

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 04:50:26