sklearn中如何将混淆矩阵坐标轴数值替换为对应字母标签
解决方案
sklearn的ConfusionMatrixDisplay原生支持自定义坐标轴标签,不需要手动修改matplotlib刻度属性,按以下步骤修改即可:
- 先修正原映射字典的语法错误:
'P': 18后缺失逗号,补全后反转字典得到「数字编码→字母标签」的映射关系 - 按类别编码升序(0到25)排序,提取出和混淆矩阵默认刻度顺序完全匹配的字母标签列表
- 初始化
ConfusionMatrixDisplay时传入display_labels参数,值为排序后的标签列表
修改后的可运行完整代码:
import matplotlib.pyplot as plt from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay # 补全语法错误后的字母-数字编码映射字典 label_map = { 'A': 7,'B': 6,'C': 14,'D': 2,'E': 19,'F': 13, 'G': 4,'H': 15, 'I': 1, 'J': 8, 'K': 24, 'L': 17, 'M': 9, 'N': 3, 'O': 11, 'P': 18, 'Q': 22, 'R': 12, 'S': 5, 'T': 0, 'U': 23, 'V': 20, 'W': 16, 'X': 10, 'Y': 21, 'Z': 25 } # 按数字编码升序排列,得到和混淆矩阵刻度顺序一致的标签列表 sorted_class_labels = [char for char, code in sorted(label_map.items(), key=lambda x: x[1])] # 计算混淆矩阵 cm = confusion_matrix(y_class, y_pred_class) # 初始化绘图对象时直接传入自定义标签 disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=sorted_class_labels) fig, ax = plt.subplots(figsize=(10,10)) ax.set_title("Confusion Matrix for Artificial neural network") disp.plot(ax=ax) plt.show()
备选方案:如果你已经完成绘图逻辑再需要修改标签,可以在
disp.plot(ax=ax)代码后追加两行代码手动替换刻度,效果和传参一致:ax.set_xticklabels(sorted_class_labels) ax.set_yticklabels(sorted_class_labels)优先推荐初始化传参的写法,能避免手动改刻度容易出现的标签错位问题。
内容的提问来源于stack exchange,提问作者Socka
相关产品推荐
相关产品推荐

