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

在Keras中使用validation_data时提取分类概率值

提取验证集概率与类别预测结果的回调实现方案

你遇到的这个问题确实很普遍——用model.fit()自带的validation_data时,Keras默认不会留存验证集的预测结果,不过自定义回调函数就能轻松解决这个需求。我给你写一个实用的回调类,训练过程中就能自动记录每轮验证集的概率值和类别预测:

from tensorflow.keras.callbacks import Callback
import numpy as np

class ValidationPredictionsCallback(Callback):
    def __init__(self, validation_data):
        super().__init__()
        self.validation_data = validation_data
        self.val_probs = []  # 存储每轮验证集的概率输出
        self.val_preds = []  # 存储每轮验证集的类别预测结果

    def on_epoch_end(self, epoch, logs=None):
        # 每轮训练结束后,对验证集执行预测
        X_val, y_val = self.validation_data
        probs = self.model.predict(X_val, verbose=0)
        # 根据任务类型调整类别预测逻辑:多分类用argmax,二分类用阈值判断
        preds = np.argmax(probs, axis=1)  # 多分类场景
        # preds = (probs > 0.5).astype(int)  # 二分类场景,可自行调整阈值
        
        self.val_probs.append(probs)
        self.val_preds.append(preds)
        print(f"第 {epoch+1} 轮验证集预测结果已保存")

使用步骤

  1. 初始化回调并传入你的验证集数据:
val_pred_callback = ValidationPredictionsCallback(validation_data=(Xtest, ytest))
  1. 将回调加入model.fit()的参数列表:
model.fit(Xtrain, ytrain, 
          epochs=epochs, 
          verbose=2, 
          batch_size=1000, 
          shuffle=True, 
          validation_data=(Xtest, ytest),
          callbacks=[val_pred_callback])

获取结果

训练结束后,你可以直接从回调实例中提取数据:

  • 所有轮次的验证集概率:val_pred_callback.val_probs(列表结构,每个元素对应一轮的概率数组)
  • 所有轮次的验证集类别预测:val_pred_callback.val_preds(同上,每个元素对应一轮的类别数组)
  • 仅获取最后一轮结果:last_epoch_probs = val_pred_callback.val_probs[-1]、last_epoch_preds = val_pred_callback.val_preds[-1]

额外提示

  • 如果验证集数据量较大,可在predict时指定batch_size参数避免内存溢出,比如self.model.predict(X_val, batch_size=1000, verbose=0)
  • 若无需保存每轮结果,只需最后一轮,可把append改为直接赋值(self.val_probs = probs),节省内存占用

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 09:50:10