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

Keras训练时如何在每个Epoch后获取验证集详细分类结果?

解决方案:自定义Callback实现Epoch级验证细节输出

嘿,这个需求我之前也碰到过,Keras本身没有直接内置一键输出验证集所有样本详细预测结果的功能,但我们可以用**自定义回调函数(Callback)**来实现每个Epoch结束后自动统计验证集里哪些样本分类正确、哪些错误,甚至保存每个样本的置信度等信息,完全不用手动跑model.predict()。

核心思路

Keras的回调函数可以在训练的关键节点(比如每个Epoch结束后)触发自定义逻辑,我们只需要重写on_epoch_end方法,在这个方法里对整个验证集做预测,然后和真实标签一一对比,就能得到所有样本的详细结果。

具体实现步骤

1. 准备验证集生成器并获取真实标签信息

首先初始化验证集的ImageDataGenerator和生成器,关键要把shuffle设为False,这样生成的样本顺序和filenames、classes属性的顺序完全对应,避免标签和预测结果错位:

from keras.preprocessing.image import ImageDataGenerator
from keras.models import Sequential
from keras.layers import Dense, Flatten
from keras.callbacks import Callback
import numpy as np
import json

# 初始化验证集数据生成器
val_datagen = ImageDataGenerator(rescale=1./255)
val_generator = val_datagen.flow_from_directory(
    'path/to/your/validation_dir',  # 替换成你的验证集路径
    target_size=(224, 224),  # 替换成你的模型输入尺寸
    batch_size=32,
    class_mode='categorical',
    shuffle=False  # 必须设为False,保证样本顺序一致
)

# 获取验证集的基础信息
val_filenames = val_generator.filenames  # 所有验证样本的文件名
val_true_labels = val_generator.classes  # 所有样本的真实标签索引
class_indices = val_generator.class_indices  # 类别名到索引的映射
idx_to_class = {v: k for k, v in class_indices.items()}  # 反转得到索引到类别名的映射

2. 编写自定义回调类

继承keras.callbacks.Callback,重写on_epoch_end方法,在每个Epoch结束后执行预测和结果统计:

class ValidationDetailedCallback(Callback):
    def __init__(self, val_gen, filenames, true_labels, idx_to_class):
        super().__init__()
        self.val_gen = val_gen
        self.filenames = filenames
        self.true_labels = true_labels
        self.idx_to_class = idx_to_class

    def on_epoch_end(self, epoch, logs=None):
        # 对整个验证集做预测
        val_pred_probs = self.model.predict(self.val_gen, verbose=0)
        val_pred_labels = np.argmax(val_pred_probs, axis=1)

        # 分类正确和错误的样本
        correct_samples = []
        incorrect_samples = []

        # 遍历所有样本,对比真实标签和预测结果
        for filename, true_idx, pred_idx, prob in zip(
            self.filenames, self.true_labels, val_pred_labels, val_pred_probs
        ):
            true_class = self.idx_to_class[true_idx]
            pred_class = self.idx_to_class[pred_idx]
            max_confidence = np.max(prob)
            
            sample_info = {
                "filename": filename,
                "true_class": true_class,
                "pred_class": pred_class,
                "confidence": round(max_confidence, 4)
            }
            
            if true_idx == pred_idx:
                correct_samples.append(sample_info)
            else:
                # 错误样本额外保存所有类别的置信度
                sample_info["all_class_probs"] = {
                    self.idx_to_class[i]: round(prob[i], 4) for i in range(len(prob))
                }
                incorrect_samples.append(sample_info)

        # 打印本次Epoch的验证细节统计
        total_samples = len(self.filenames)
        correct_rate = len(correct_samples) / total_samples * 100
        incorrect_rate = len(incorrect_samples) / total_samples * 100

        print(f"\n=== Epoch {epoch+1} 验证集详细结果 ===")
        print(f"总样本数: {total_samples}")
        print(f"分类正确: {len(correct_samples)} ({correct_rate:.2f}%)")
        print(f"分类错误: {len(incorrect_samples)} ({incorrect_rate:.2f}%)")

        # 打印前5个错误样本示例(可根据需求调整)
        if incorrect_samples:
            print("\n错误样本示例:")
            for idx, sample in enumerate(incorrect_samples[:5], 1):
                print(f"{idx}. 文件: {sample['filename']} | 真实类别: {sample['true_class']} | 预测类别: {sample['pred_class']} | 置信度: {sample['confidence']}")

        # 将结果保存到JSON文件,方便后续分析
        with open(f"epoch_{epoch+1}_validation_results.json", "w", encoding="utf-8") as f:
            json.dump({
                "epoch": epoch+1,
                "total_samples": total_samples,
                "correct_count": len(correct_samples),
                "incorrect_count": len(incorrect_samples),
                "correct_samples": correct_samples,
                "incorrect_samples": incorrect_samples
            }, f, indent=4, ensure_ascii=False)

3. 训练时加入自定义回调

在fit_generator(或者新版Keras的fit)中,把这个回调加入到callbacks列表里即可:

# 假设你已经定义好了自己的模型
model = Sequential([
    # 替换成你的模型层,比如卷积层、池化层等
    Flatten(),
    Dense(5, activation="softmax")  # 5类对应输出5个神经元
])

model.compile(
    optimizer="adam",
    loss="categorical_crossentropy",
    metrics=["accuracy"]
)

# 初始化自定义回调
val_details_callback = ValidationDetailedCallback(
    val_gen=val_generator,
    filenames=val_filenames,
    true_labels=val_true_labels,
    idx_to_class=idx_to_class
)

# 开始训练
model.fit_generator(
    train_generator,  # 替换成你的训练集生成器
    epochs=10,
    validation_data=val_generator,
    callbacks=[val_details_callback]
)

注意事项

  • shuffle=False很重要:如果验证集生成器开启了shuffle,预测结果的顺序会和filenames、true_labels不匹配,导致标签对应错误。
  • 性能考量:如果验证集非常大,predict过程可能会耗时,可以适当调大验证集生成器的batch_size来加速。
  • 新版Keras兼容:如果你用的是TensorFlow 2.x集成的Keras,fit_generator已经被fit替代,直接用model.fit(train_generator, ...)即可,回调的用法完全一致。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 09:07:39