使用flow_from_directory构建图像数据集训练模型后,如何获取混淆矩阵并绘制ROC曲线与混淆矩阵
如何在flow_from_directory场景下绘制混淆矩阵与ROC曲线?
嘿,我来帮你搞定这个问题!你已经用ImageDataGenerator完成了模型训练,现在想生成混淆矩阵和ROC曲线对吧?这在flow_from_directory的场景下完全可以实现,我一步步给你讲清楚:
一、生成混淆矩阵
要绘制混淆矩阵,核心是获取测试集的真实标签和模型的预测类别,然后用sklearn的工具来生成和可视化。
关键前置提醒
首先,在创建测试集(或者验证集)的时候,一定要设置shuffle=False!因为flow_from_directory默认会打乱数据顺序,这会导致真实标签和预测结果的索引不匹配,混淆矩阵就会完全出错。修改你的test_dataset代码:
test_dataset = test_datagen.flow_from_directory( directory = './test', target_size = tsize, class_mode = 'categorical', batch_size = BS, shuffle=False # 必须添加这个参数! )
完整代码实现
import numpy as np import matplotlib.pyplot as plt from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay # 1. 获取测试集的真实标签 y_true = test_dataset.classes # 2. 获取模型对测试集的预测概率,再转换为预测类别(取概率最大的类别) y_pred_probs = model.predict(test_dataset, verbose=1) y_pred = np.argmax(y_pred_probs, axis=1) # 3. 获取类别名称(从数据集的class_indices字典中提取) class_names = list(test_dataset.class_indices.keys()) # 4. 生成并绘制混淆矩阵 confusion_mat = confusion_matrix(y_true, y_pred) disp = ConfusionMatrixDisplay( confusion_matrix=confusion_mat, display_labels=class_names ) disp.plot(cmap=plt.cm.Blues) plt.title("Test Set Confusion Matrix") plt.show()
二、绘制ROC曲线
ROC曲线需要用到真实标签的one-hot编码和模型对每个类别的预测概率。对于多分类任务,我们可以绘制每个类别的ROC曲线,或者计算宏平均/微平均的ROC曲线。
完整代码实现
from sklearn.metrics import roc_curve, auc from sklearn.preprocessing import label_binarize # 1. 将真实标签转换为one-hot编码格式 y_true_onehot = label_binarize(y_true, classes=np.arange(len(class_names))) n_classes = y_true_onehot.shape[1] # 2. 逐个类别计算ROC曲线的FPR、TPR和AUC值 fpr = dict() tpr = dict() roc_auc = dict() for i in range(n_classes): fpr[i], tpr[i], _ = roc_curve(y_true_onehot[:, i], y_pred_probs[:, i]) roc_auc[i] = auc(fpr[i], tpr[i]) # 3. 绘制多分类ROC曲线 plt.figure(figsize=(8, 6)) # 可以根据类别数量调整颜色列表 colors = ['blue', 'red', 'green', 'orange', 'purple'] for i, color in zip(range(n_classes), colors): plt.plot( fpr[i], tpr[i], color=color, lw=2, label=f'ROC curve of {class_names[i]} (AUC = {roc_auc[i]:.2f})' ) # 绘制随机猜测的基准线 plt.plot([0, 1], [0, 1], 'k--', lw=2) plt.xlim([0.0, 1.0]) plt.ylim([0.0, 1.05]) plt.xlabel('False Positive Rate') plt.ylabel('True Positive Rate') plt.title('Multi-class ROC Curve') plt.legend(loc="lower right") plt.show()
如果需要给验证集生成混淆矩阵或ROC曲线,只需要把上述代码中的test_dataset换成valid_dataset,同样记得给valid_dataset加上shuffle=False参数哦!
内容的提问来源于stack exchange,提问作者Carlos Berrocal
相关产品推荐
相关产品推荐

