scikit-learn ConfusionMatrixDisplay无预测样本类时标签偏移问题及解决
问题描述
我有一个可预测10个类别的模型,预测了26个样本,但y_test和y_pred中都没有class 3('jumping_jacks')的样本。原本预期混淆矩阵中“真实标签jumping_jacks”行和“预测标签jumping_jacks”列全为0,但实际显示的class 3对应的是class 4('lateral_shoulder_raises')的预测结果,从第3行/列开始所有内容都偏移,导致包含样本的class 9('tricep_extensions')未在矩阵中显示。
原输出问题示例

复现代码
from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay from sklearn.preprocessing import LabelEncoder import matplotlib.pyplot as plt import pandas as pd import numpy as np ex_classes = {'Classes': ['bicep_curls', 'dumbbell_rows', 'dumbbell_shoulder_press', 'jumping_jacks', 'lateral_shoulder_raises', 'lunges', 'pushups', 'situps', 'squats', 'tricep_extensions']} df_classes = pd.DataFrame(data=ex_classes) label_enc = LabelEncoder() label_enc.fit(df_classes['Classes']) y_test = np.asarray([8, 8, 8, 6, 6, 6, 2, 2, 2, 5, 5, 5, 1, 1, 1, 7, 7, 7, 9, 9, 9, 0, 0, 0, 0, 4]) y_pred = np.asarray([8, 4, 4, 6, 6, 6, 2, 2, 2, 5, 5, 5, 1, 1, 1, 9, 7, 7, 9, 9, 9, 0, 0, 1, 0, 4]) cm = confusion_matrix(y_test, y_pred) display = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels = label_enc.classes_) fig, ax = plt.subplots(figsize=(10,10)) display.plot(ax=ax, xticks_rotation='vertical') plt.show()
问题原因
confusion_matrix默认只会统计y_test和y_pred中实际出现的类别,不会自动包含训练时定义的所有类别。你的标签编码器对应10个类别,但测试和预测结果里仅包含9个(缺少class 3),所以生成的混淆矩阵是9x9维度。但你给ConfusionMatrixDisplay传入了10个类别的标签列表,导致标签和矩阵内容错位,出现偏移和类别缺失。
解决方法
要强制混淆矩阵包含所有预设类别,只需在计算混淆矩阵时指定labels参数为所有类别的编码值:
- 获取所有类别的编码值:
all_classes = label_enc.transform(label_enc.classes_) - 计算混淆矩阵时传入
labels=all_classes,确保矩阵维度与类别总数一致 - 再用
ConfusionMatrixDisplay可视化,此时标签与矩阵内容能一一对应
修改后的完整代码
from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay from sklearn.preprocessing import LabelEncoder import matplotlib.pyplot as plt import pandas as pd import numpy as np ex_classes = {'Classes': ['bicep_curls', 'dumbbell_rows', 'dumbbell_shoulder_press', 'jumping_jacks', 'lateral_shoulder_raises', 'lunges', 'pushups', 'situps', 'squats', 'tricep_extensions']} df_classes = pd.DataFrame(data=ex_classes) label_enc = LabelEncoder() label_enc.fit(df_classes['Classes']) # 获取所有类别的编码值,确保包含全部10个类别 all_classes = label_enc.transform(label_enc.classes_) y_test = np.asarray([8, 8, 8, 6, 6, 6, 2, 2, 2, 5, 5, 5, 1, 1, 1, 7, 7, 7, 9, 9, 9, 0, 0, 0, 0, 4]) y_pred = np.asarray([8, 4, 4, 6, 6, 6, 2, 2, 2, 5, 5, 5, 1, 1, 1, 9, 7, 7, 9, 9, 9, 0, 0, 1, 0, 4]) # 计算混淆矩阵时指定labels参数,强制包含所有类别 cm = confusion_matrix(y_test, y_pred, labels=all_classes) display = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=label_enc.classes_) fig, ax = plt.subplots(figsize=(10,10)) display.plot(ax=ax, xticks_rotation='vertical') plt.show()
效果说明
修改后生成的混淆矩阵为10x10维度,其中'jumping_jacks'对应的行和列全为0,所有10个类别都会正常显示,不会再出现偏移或类别缺失的问题。
内容的提问来源于stack exchange,提问作者lonyen11
相关产品推荐
相关产品推荐

