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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 03:06:09