在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} 轮验证集预测结果已保存")
使用步骤
- 初始化回调并传入你的验证集数据:
val_pred_callback = ValidationPredictionsCallback(validation_data=(Xtest, ytest))
- 将回调加入
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
相关产品推荐
相关产品推荐

