绘制3分类混淆矩阵时遇FixedLocator标签不匹配错误,如何修复?
修复sklearn混淆矩阵维度不匹配与标签错误问题
错误根源
- 标签格式不匹配:代码中注释掉了
testY=np.argmax(testY, axis=1),如果testY是独热编码格式(每个样本是数组),直接传入confusion_matrix会导致函数无法正确识别类别;如果testY是1/2/3的标签,但y_prediction是argmax后得到的0/1/2索引,两者的类别集合合并后会包含0/1/2/3四个值,因此生成4x4的混淆矩阵。 - 未指定明确类别范围:
confusion_matrix默认会自动根据输入的真实标签和预测标签的所有唯一值生成类别列表,当出现多余的类别(比如0)时,矩阵维度超出预期,和指定的3个display_labels不匹配,触发FixedLocator错误。
修复方案
方案一:统一使用0/1/2作为类别索引(推荐,适配模型argmax输出)
取消注释testY的argmax转换,确保真实标签和预测标签都是0/1/2的索引格式,再对应设置显示标签:
from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay import numpy as np import matplotlib.pyplot as plt # 关键:将独热编码的testY转换为类别索引 testY = np.argmax(testY, axis=1) y_prediction = model.predict(testX) y_prediction = np.argmax(y_prediction, axis=1) # 定义对应0/1/2的显示标签 labels = ['1', '2', '3'] # 生成归一化混淆矩阵 result = confusion_matrix(testY, y_prediction, normalize='pred') print(result) # 生成混淆矩阵并可视化 cm = confusion_matrix(testY, y_prediction) disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=labels) disp.plot(cmap=plt.cm.Blues) plt.show()
方案二:保留1/2/3作为类别标签
如果需要保持标签为1/2/3,需要给confusion_matrix明确指定labels参数,限定类别范围:
from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay import numpy as np import matplotlib.pyplot as plt # 将独热编码转换为1/2/3的标签 testY = np.argmax(testY, axis=1) + 1 y_prediction = model.predict(testX) y_prediction = np.argmax(y_prediction, axis=1) + 1 # 明确指定类别范围为[1,2,3] target_labels = [1, 2, 3] display_labels = ['1', '2', '3'] result = confusion_matrix(testY, y_prediction, labels=target_labels, normalize='pred') print(result) cm = confusion_matrix(testY, y_prediction, labels=target_labels) disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=display_labels) disp.plot(cmap=plt.cm.Blues) plt.show()
验证效果
修复后,混淆矩阵会变成3x3的格式,不再出现首行全0的情况,同时可视化时标签数量与刻度数量匹配,不会触发ValueError。
内容的提问来源于stack exchange,提问作者Colab4066
相关产品推荐
相关产品推荐

