使用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
相关产品推荐
相关产品推荐

