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

TensorFlow训练中用validation_split后如何获取验证数据绘制ROC曲线

解决方法

当你用validation_split时,Keras的History对象不会保存验证集的原始数据,所以直接从history.validation_data拿数据是行不通的。最稳妥的方式是手动拆分数据集,这样能明确获取到验证集的特征和标签,方便后续绘制ROC曲线。

步骤1:手动拆分训练集和验证集

用sklearn.model_selection.train_test_split来拆分数据,代替validation_split参数,这样你能直接拿到X_val和y_val:

from sklearn.model_selection import train_test_split
import keras
from sklearn.metrics import roc_curve, auc
from sklearn.preprocessing import label_binarize
import matplotlib.pyplot as plt
import numpy as np

# 手动拆分数据,测试集比例设为0.15,和原来的validation_split一致
train_X_split, X_val, train_y_split, y_val = train_test_split(
    train_X, train_y, test_size=0.15, random_state=42  # random_state固定拆分结果,可选
)

# 模型编译(和你原来的代码一致)
loss = keras.losses.categorical_crossentropy
optim = keras.optimizers.Adam(learning_rate=0.0009)
metrics = ["accuracy"]
model_lstm.compile(loss=loss, optimizer=optim, metrics=metrics)

# 训练时用validation_data传入手动拆分的验证集
history = model_lstm.fit(
    train_X_split, train_y_split,
    batch_size=32, epochs=10,
    validation_data=(X_val, y_val),
    callbacks=CALLBACKS
)

步骤2:预测验证集并绘制ROC曲线

因为你用的是categorical_crossentropy,属于多分类任务,所以需要先把标签二值化,再计算每个类别的ROC曲线:

# 预测验证集的概率
y_pred = model_lstm.predict(X_val)

# 获取类别数量
n_classes = train_y.shape[1]

# 将真实标签二值化(如果已经是one-hot编码可以跳过这步)
y_val_bin = label_binarize(y_val, classes=np.arange(n_classes))

# 计算每个类别的ROC曲线和AUC值
fpr = dict()
tpr = dict()
roc_auc = dict()
for i in range(n_classes):
    fpr[i], tpr[i], _ = roc_curve(y_val_bin[:, i], y_pred[:, i])
    roc_auc[i] = auc(fpr[i], tpr[i])

# 绘制所有类别的ROC曲线
plt.figure()
colors = ['blue', 'red', 'green', 'orange']  # 根据类别数量调整颜色
for i, color in zip(range(n_classes), colors):
    plt.plot(fpr[i], tpr[i], color=color, lw=2,
             label='ROC curve of class {0} (area = {1:0.2f})'
             ''.format(i, roc_auc[i]))

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_curve(y_val, y_pred[:,1])即可(假设正类是第二个类别)。
  • 手动拆分数据的好处是你能完全掌控验证集的内容,后续做其他分析也更方便,避免依赖Keras内部的拆分逻辑。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 21:50:22