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

如何基于多参数使用Keras ModelCheckpoint保存最优模型

关于Keras ModelCheckpoint多指标监控的解决方案

好问题!原生的Keras ModelCheckpoint回调确实不支持直接给monitor参数设置多个指标——它只能基于单一指标的变化来触发模型保存逻辑。不过你的需求(优先看验证准确率,相同情况下选训练准确率更高的模型)完全可以实现,下面分享两种实用方案:

方案一:自定义回调函数(最贴合你的需求)

既然原生回调满足不了多条件判断,咱可以自己写一个回调类,在每个epoch结束后手动实现你的保存逻辑。核心思路是:

  • 在每个epoch结束时,获取当前的验证准确率(val_accuracy)和训练准确率(accuracy)
  • 维护一个变量记录历史最优的验证准确率和对应的训练准确率
  • 按照你的规则判断是否保存模型:
    1. 如果当前验证准确率 > 历史最优 → 保存模型,并更新历史记录
    2. 如果当前验证准确率 == 历史最优,但当前训练准确率 > 历史训练准确率 → 覆盖保存模型,并更新历史记录

下面是一个简单的代码示例:

from tensorflow.keras.callbacks import Callback
import tensorflow as tf

class CustomModelCheckpoint(Callback):
    def __init__(self, filepath, monitor_val='val_accuracy', monitor_train='accuracy', mode='max'):
        super().__init__()
        self.filepath = filepath
        self.monitor_val = monitor_val
        self.monitor_train = monitor_train
        self.mode = mode
        self.best_val = -float('inf') if mode == 'max' else float('inf')
        self.best_train = -float('inf') if mode == 'max' else float('inf')

    def on_epoch_end(self, epoch, logs=None):
        logs = logs or {}
        current_val = logs.get(self.monitor_val)
        current_train = logs.get(self.monitor_train)
        
        if current_val is None or current_train is None:
            print(f"Warning: {self.monitor_val} or {self.monitor_train} not found in logs")
            return
        
        # 判断是否需要保存模型
        save = False
        if self.mode == 'max':
            if current_val > self.best_val:
                save = True
            elif current_val == self.best_val and current_train > self.best_train:
                save = True
        else:
            # 针对min模式的逻辑,你的需求用不到,但可保留扩展
            if current_val < self.best_val:
                save = True
            elif current_val == self.best_val and current_train < self.best_train:
                save = True
        
        if save:
            print(f"\nEpoch {epoch+1}: {self.monitor_val} improved from {self.best_val:.4f} to {current_val:.4f}, saving model to {self.filepath}")
            self.model.save(self.filepath)
            self.best_val = current_val
            self.best_train = current_train
        else:
            print(f"\nEpoch {epoch+1}: {self.monitor_val} did not improve or train accuracy is not better")

# 使用示例
checkpoint = CustomModelCheckpoint(filepath='best_model.h5')
model.fit(..., callbacks=[checkpoint])

方案二:用调和平均合并两个指标

如果更倾向于复用原生ModelCheckpoint的功能,可以把验证准确率和训练准确率通过调和平均合并成一个单一指标,然后让ModelCheckpoint监控这个合成指标。

调和平均非常适合处理准确率这类比例型指标,它的公式是:
combined_score = 2 * (train_acc * val_acc) / (train_acc + val_acc)

这个分数能同时兼顾两个指标的表现,而且当其中一个指标很低时,整体分数也会被拉低,符合你优先保证验证准确率、同时兼顾训练准确率的需求。

实现方式是自定义一个回调,在每个epoch计算这个合成指标并加入到logs中,然后让ModelCheckpoint监控这个指标:

from tensorflow.keras.callbacks import Callback, ModelCheckpoint

class CombinedScoreCallback(Callback):
    def on_epoch_end(self, epoch, logs=None):
        logs = logs or {}
        train_acc = logs.get('accuracy')
        val_acc = logs.get('val_accuracy')
        if train_acc and val_acc:
            # 计算调和平均
            combined_score = 2 * (train_acc * val_acc) / (train_acc + val_acc)
            logs['combined_score'] = combined_score

# 使用示例
combined_callback = CombinedScoreCallback()
checkpoint = ModelCheckpoint(
    filepath='best_model.h5',
    monitor='combined_score',
    mode='max',
    save_best_only=True
)
model.fit(..., callbacks=[combined_callback, checkpoint])

方案对比

  • 自定义回调:完全贴合你的规则,逻辑清晰,可灵活修改判断条件(比如加入epoch限制、损失值等)
  • 调和平均:复用原生ModelCheckpoint的功能,代码更简洁,但合成指标的逻辑是固定的,不如自定义回调灵活

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 09:31:39