You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.30 13:28:14