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

如何使用ImageDataGenerator生成训练模型的ROC曲线

解决使用ImageDataGenerator生成数据时绘制ROC曲线的问题

我看到你已经给test_generator设置了shuffle=False(这一步特别关键,能保证预测结果和真实标签的顺序严格对应),但报错的核心问题是**roc_curve的参数顺序搞反了**,另外还要注意预测结果的维度处理。下面是完整的修正步骤和代码:

一、错误原因分析

你原来的代码写法:

fpr, tpr, tresholds = roc_curve(y_pred, test_generator.classes)

而roc_curve的正确参数顺序是**(y_true, y_score)**——真实标签在前,预测概率在后;同时,model.predict_generator返回的是(样本数, 1)的二维数组,需要转成一维才能和一维的test_generator.classes匹配。

二、完整的ROC曲线绘制步骤

1. 获取预测概率并调整维度

先执行预测,再把二维的预测结果展平为一维数组:

# 获取模型对测试集的预测概率
y_pred = model.predict_generator(test_generator, steps=step_size_test)
# 将二维数组展平为一维,适配真实标签的维度
y_pred = y_pred.ravel()

2. 获取真实标签

因为你设置了shuffle=False,所以test_generator.classes的顺序和预测结果完全对应:

y_true = test_generator.classes

3. 计算ROC曲线参数与AUC值

from sklearn.metrics import roc_curve, roc_auc_score

# 计算假阳性率(fpr)、真阳性率(tpr)以及对应的阈值
fpr, tpr, thresholds = roc_curve(y_true, y_pred)
# 计算AUC(曲线下面积)值
auc_score = roc_auc_score(y_true, y_pred)

4. 绘制ROC曲线

用matplotlib完成可视化:

import matplotlib.pyplot as plt

plt.figure(figsize=(8,6))
# 绘制ROC曲线
plt.plot(fpr, tpr, label=f'ROC Curve (AUC = {auc_score:.4f})')
# 绘制随机猜测的基准线
plt.plot([0,1], [0,1], 'k--')
plt.xlabel('False Positive Rate')
plt.ylabel('True Positive Rate')
plt.title('Receiver Operating Characteristic (ROC) Curve')
plt.legend(loc='lower right')
plt.show()

三、额外注意事项

  • 确保模型最后一层是**Dense(1, activation='sigmoid')**:因为你用的是binary_crossentropy损失函数,只有sigmoid激活才能输出0-1之间的概率值,这是计算ROC曲线的必要前提。
  • 如果测试集样本数不能被batch_size整除,你可以把steps设为step_size_test + 1,之后用y_pred[:test_generator.n]截取前test_generator.n个结果,避免多余的预测值干扰计算。

内容的提问来源于stack exchange,提问作者Leslie

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.12 05:05:42