如何为Sklearn生成的confusion matrix正确标注标签?
解决Sklearn混淆矩阵的标签匹配问题
嗨,这个问题太常见了——当类别数量多或者部分类别在测试集中没出现时,手动硬编码标签很容易搞混顺序。好在Sklearn已经提供了准确获取混淆矩阵对应标签顺序的方法,完全不用自己瞎猜!
核心解决方法:用unique_labels获取正确标签顺序
Sklearn生成的混淆矩阵,其行对应真实标签,列对应预测标签,而这些标签的顺序是真实标签与预测标签的并集按字典序排序后的结果。你可以直接使用sklearn.utils.multiclass.unique_labels函数来获取这个顺序,它会自动识别所有出现在y_test和y_predicted_counts中的标签,确保和矩阵的行、列完全对应。
具体修改你的代码
首先,导入unique_labels工具:
from sklearn.utils.multiclass import unique_labels
接着,在生成混淆矩阵后,获取对应的标签列表:
cm = confusion_matrix(y_test, y_predicted_counts) # 获取与混淆矩阵行/列完全匹配的标签顺序 class_labels = unique_labels(y_test, y_predicted_counts)
最后,把绘制混淆矩阵的代码里的classes参数替换成class_labels:
plot = plot_confusion_matrix(cm, classes=class_labels, normalize=False, title='Confusion matrix')
为什么这能解决你的问题?
unique_labels(y_true, y_pred)会返回所有在真实标签和预测标签中出现的唯一标签,并且按排序后的顺序排列,这个顺序和confusion_matrix生成的矩阵的行、列顺序完全一致。- 比如你提供的13x13混淆矩阵,
class_labels会返回13个标签,其中class_labels[0]对应矩阵第0行的真实标签,class_labels[6]对应矩阵第6列的预测标签——对应你矩阵里第二行第七列的3,就代表真实标签是class_labels[1]的样本,有3个被预测成了class_labels[6]。
额外验证小技巧
你可以打印class_labels来确认顺序:
print("Confusion Matrix Label Order:", class_labels)
这样就能直观看到每个矩阵索引对应的具体标签,避免混淆。
内容的提问来源于stack exchange,提问作者mikelowry
相关产品推荐
相关产品推荐

