如何用plot_confusion_matrix在同一窗口绘制带标签的三个混淆矩阵
解决方案
要给每个混淆矩阵子图添加类别标签和轴名称,你可以通过手动设置子图的坐标轴属性来解决。以下是修改后的完整代码:
import numpy as np import matplotlib.pyplot as plt from sklearn.metrics import plot_confusion_matrix # 假设使用sklearn的该函数 # 定义类别标签 label_ax = ['0','1','2','3','4','5','6','7','8','9','10','11','12','13','14','15','16','17','18'] # 计算累计混淆矩阵 a1 = np.zeros(shape=(19, 19)) a2 = np.zeros(shape=(19, 19)) a3 = np.zeros(shape=(19, 19)) for element in total_cm_svm: a1 = a1 + element for element in total_cm_lda: a2 = a2 + element for element in total_cm_etc: a3 = a3 + element # 创建子图布局,调整画布大小避免标签拥挤 fig, axes = plt.subplots(nrows=1, ncols=3, figsize=(18, 6)) # 绘制SVM混淆矩阵并设置标签 cm = plot_confusion_matrix(conf_mat=a1, colorbar=False, class_names=label_ax, axis=axes[0]) axes[0].set_title('SVM 混淆矩阵') axes[0].set_xlabel('预测类别') axes[0].set_ylabel('真实类别') axes[0].set_xticklabels(label_ax, rotation=45) axes[0].set_yticklabels(label_ax) # 绘制LDA混淆矩阵并设置标签 cm1 = plot_confusion_matrix(conf_mat=a2, colorbar=False, class_names=label_ax, axis=axes[1]) axes[1].set_title('LDA 混淆矩阵') axes[1].set_xlabel('预测类别') axes[1].set_ylabel('真实类别') axes[1].set_xticklabels(label_ax, rotation=45) axes[1].set_yticklabels(label_ax) # 绘制ETC混淆矩阵并设置标签 cm2 = plot_confusion_matrix(conf_mat=a3, colorbar=False, class_names=label_ax, axis=axes[2]) axes[2].set_title('ETC 混淆矩阵') axes[2].set_xlabel('预测类别') axes[2].set_ylabel('真实类别') axes[2].set_xticklabels(label_ax, rotation=45) axes[2].set_yticklabels(label_ax) # 自动调整子图间距,防止标签被截断 plt.tight_layout() plt.show()
关键修改说明:
- 显式定义类别标签:确保
label_ax变量指向你需要的19个类别标签列表,统一应用到所有子图。 - 设置轴名称:通过
set_xlabel和set_ylabel分别为每个子图添加"预测类别"和"真实类别"的轴标签。 - 强制设置刻度标签:调用
set_xticklabels和set_yticklabels手动绑定类别标签到坐标轴,配合rotation=45避免x轴标签重叠。 - 优化布局:调整
figsize增大画布宽度,并用plt.tight_layout()自动调整子图间距,防止标签被截断或遮挡。 - 添加子图标题:给每个子图添加模型名称标题,方便区分不同混淆矩阵对应的模型。
内容的提问来源于stack exchange,提问作者Renlo
相关产品推荐
相关产品推荐

