如何在Scikit-learn中为混淆矩阵打印类别标签?
为Scikit-learn混淆矩阵添加类别标签并优化排版
我完全理解你的困扰——13类的混淆矩阵没有标签,根本没法对应到具体类别,排版还乱糟糟的,完全看不出模型在哪些类别上表现好。别担心,有几个简单的方法可以解决这个问题,让混淆矩阵清晰易懂:
方法1:用Pandas把混淆矩阵转成带标签的DataFrame
这是最直接的方式,能让矩阵自带行和列的类别标签,排版也整齐。
步骤如下:
- 先获取排序后的类别列表(确保和混淆矩阵的行列顺序一致)
- 生成混淆矩阵时显式指定
labels参数 - 把矩阵转成Pandas DataFrame,设置行索引和列名为类别标签
修改你的代码如下:
import pandas as pd # 别忘了导入Pandas def classify_data(df, feature_cols, file): nbr_folds = 5 RANDOM_STATE = 0 attributes = df.loc[:, feature_cols] class_label = df['task'] # 获取排序后的类别标签,确保顺序统一 class_labels = sorted(class_label.unique()) file.write("\nFeatures used: ") for feature in feature_cols: file.write(feature + ",") print("Features used", feature_cols) sampler = RandomOverSampler(random_state=RANDOM_STATE) print("RandomForest") file.write("\nRandomForest") rfc = RandomForestClassifier(max_depth=2, random_state=RANDOM_STATE) pipeline = make_pipeline(sampler, rfc) class_label_predicted = cross_val_predict(pipeline, attributes, class_label, cv=nbr_folds) # 生成混淆矩阵时指定labels参数,确保行列顺序和我们的类别列表一致 conf_mat = confusion_matrix(class_label, class_label_predicted, labels=class_labels) # 转成带标签的DataFrame,明确标注真实/预测类别 conf_mat_df = pd.DataFrame(conf_mat, index=[f"True: {cls}" for cls in class_labels], columns=[f"Pred: {cls}" for cls in class_labels]) print("\nConfusion Matrix with Labels:") print(conf_mat_df) accuracy = accuracy_score(class_label, class_label_predicted) print("Rows classified: " + str(len(class_label_predicted))) print("Accuracy: {0:.3f}%\n".format(accuracy * 100)) file.write("\nClassifier settings:" + str(pipeline) + "\n") file.write("\nRows classified: " + str(len(class_label_predicted))) file.write("\nAccuracy: {0:.3f}%\n".format(accuracy * 100)) # 把带标签的矩阵写入文件 file.write("\nConfusion Matrix with Labels:\n") file.write(conf_mat_df.to_string()) file.write("\n")
这样打印出来的矩阵会清晰显示每一行是真实类别,每一列是预测类别,完全不会混淆。
方法2:用Scikit-learn自带的可视化工具
如果想要更直观的可视化效果,可以用ConfusionMatrixDisplay生成带标签的混淆矩阵图:
from sklearn.metrics import ConfusionMatrixDisplay import matplotlib.pyplot as plt # 生成可视化对象 disp = ConfusionMatrixDisplay(confusion_matrix=conf_mat, display_labels=class_labels) # 调整字体大小,避免标签重叠 disp.plot(cmap="Blues", xticks_rotation=45) plt.tight_layout() plt.show()
这个图会自动把类别标签标在坐标轴上,还能通过颜色深浅直观看到分类的对错情况,非常适合快速分析模型表现。
为什么之前的矩阵没有标签?
默认情况下,confusion_matrix会按类别在数据中出现的顺序(或者排序后的顺序)生成矩阵,但不会显式标注类别名称。通过指定labels参数,我们可以固定矩阵的行列顺序,再结合DataFrame或可视化工具,就能把类别标签对应上了。
内容的提问来源于stack exchange,提问作者fall2
相关产品推荐
相关产品推荐

