Keras模型训练完成后如何逐epoch计算分类评估指标
TensorFlow Keras 逐epoch计算全量分类指标方案
要获取单个epoch对应的模型状态、计算包含各类别precision/recall/f1的完整分类报告,有两种可直接落地的实现方式,分别适配「未启动训练/可重跑训练」和「已经完成训练」两种场景。
方案1:自定义回调(训练时实时计算,效率最高)
核心逻辑是继承Keras Callback基类,重写on_epoch_end钩子方法,在每个epoch训练结束时,直接用当前内存中已更新完权重的模型做验证集预测,计算并输出分类报告,无需事后加载模型文件。
依赖导入
import numpy as np import glob from pathlib import Path import tensorflow as tf from tensorflow import keras from sklearn.metrics import classification_report
自定义回调实现
class EpochClsReportCallback(tf.keras.callbacks.Callback): def __init__(self, val_x, val_y, class_names=None): super().__init__() self.val_x = val_x # 自动适配one-hot标签和整数ID标签 self.val_y_true = np.argmax(val_y, axis=-1) if len(val_y.shape) > 1 else val_y self.class_names = class_names def on_epoch_end(self, epoch, logs=None): # 用当前epoch更新完权重的模型做预测 y_pred = self.model.predict(self.val_x, verbose=0) y_pred_id = np.argmax(y_pred, axis=-1) # 打印当前epoch指标 print(f"\n================ Epoch {epoch+1} 验证集分类指标 ================") print(classification_report( y_true=self.val_y_true, y_pred=y_pred_id, target_names=self.class_names, digits=2 ))
接入原有训练代码
只需要把自定义回调加入原有CALLBACKS列表即可,其他编译、训练逻辑无需改动:
# 初始化回调,100分类场景下类别名直接用0-99的字符串即可 eval_cb = EpochClsReportCallback( val_x=test_X, val_y=test_y, class_names=[str(i) for i in range(100)] ) CALLBACKS = [ tf.keras.callbacks.ModelCheckpoint( filepath=Path(logpath, 'model_checkpoint-{epoch:02d}-{val_loss:.2f}.h5'), verbose=1, save_weights_only=False, save_freq='epoch' ), tensorboard, eval_cb # 新增自定义回调 ] loss = keras.losses.categorical_crossentropy optim = keras.optimizers.Adam(learning_rate=0.0009) metrics = ["accuracy"] model.compile(loss=loss, optimizer=optim, metrics=metrics) history = model.fit( train_X, train_y, batch_size=32, epochs=10, validation_data=(test_X, test_y), callbacks=CALLBACKS )
方案2:加载已保存的checkpoint计算(适配已完成训练的场景)
如果你已经跑完训练流程、没有提前加上述回调,可以直接遍历现有代码保存的所有epoch checkpoint文件,逐个加载模型后计算指标,不需要重跑训练:
# 自动适配one-hot标签和整数ID标签 y_true = np.argmax(test_y, axis=-1) if len(test_y.shape) > 1 else test_y # 遍历所有保存的checkpoint,按epoch编号排序 for ckpt_path in sorted(glob.glob(str(Path(logpath, "model_checkpoint-*.h5")))): # 从文件名提取epoch编号 epoch_id = int(ckpt_path.split("-")[1]) # 加载对应epoch的模型 model = keras.models.load_model(ckpt_path) # 预测计算指标 y_pred = model.predict(test_X, verbose=0) y_pred_id = np.argmax(y_pred, axis=-1) print(f"\n================ Checkpoint Epoch {epoch_id} 验证集分类指标 ================") print(classification_report( y_true=y_true, y_pred=y_pred_id, target_names=[str(i) for i in range(100)], digits=2 ))
注意事项
- 两种方案输出的报告格式完全匹配需求:逐行展示每个类别的precision、recall、f1-score、support,末尾自动生成整体准确率、macro avg、weighted avg三类汇总统计值
- 回调方案直接读取内存中的模型权重,没有磁盘IO开销,计算效率远高于加载checkpoint的方案
- 如果需要持久化存储每个epoch的指标,可以将
classification_report参数设置为output_dict=True,将返回的指标字典追加到列表中,训练/遍历结束后可导出为CSV等格式 - 如果验证集数据量较大,可以在
predict时设置合适的batch_size参数,避免显存溢出。
内容的提问来源于stack exchange,提问作者Test
相关产品推荐
相关产品推荐

