如何按对角线值降序对pandas生成的混淆矩阵排序
混淆矩阵按对角线值降序重排实现
完全可以实现,核心逻辑是先提取每个类别对应对角线位置的预测正确样本数,按数值从大到小排序得到新的类别顺序,再同步重排混淆矩阵的行和列即可。
在你原有代码基础上补充重排逻辑即可,完整代码如下:
import pandas as pd import seaborn as sn import matplotlib.pyplot as plt data = {'y_Actual': [3, 3, 1, 1, 0, 1, 2, 3, 1, 1, 1, 0, 2, 4, 3], 'y_Predicted': [1, 2, 2, 1, 0, 1, 3, 0, 1, 0, 0, 0, 3, 4, 2] } df = pd.DataFrame(data, columns=['y_Actual','y_Predicted']) confusion_matrix = pd.crosstab(df['y_Actual'], df['y_Predicted'], rownames=['Actual'], colnames=['Predicted']) # 补全可能缺失的类别行列,避免维度不匹配,所有类别都覆盖时可省略 confusion_matrix = confusion_matrix.reindex(index=range(5), columns=range(5), fill_value=0) # 提取对角线值,生成降序排列的类别索引 sorted_idx = confusion_matrix.values.diagonal().argsort()[::-1] # 用相同索引同步重排行、列,保证类别对应关系正确 sorted_confusion_matrix = confusion_matrix.iloc[sorted_idx, sorted_idx] sn.heatmap(sorted_confusion_matrix, annot=True) plt.show()
关键注意点
- 重排行和列必须使用完全一致的索引顺序,否则会打乱真实类别和预测类别的对应关系,导致混淆矩阵统计错误。
- 上述示例里
range(5)对应0-4共5个分类类别,实际使用时替换成你自己任务的全量类别列表即可。 - 你当前示例的原始混淆矩阵效果如下:

运行重排代码后,对角线会从左上角开始按数值从高到低排列,预测正确样本数最多的类别会排在最靠前的位置。
内容的提问来源于stack exchange,提问作者codelifevcd
相关产品推荐
相关产品推荐

