如何在Precision与Recall均更优时保存Keras模型?
实现Precision和Recall双指标均更优时保存Keras模型
你说得对,直接加两个ModelCheckpoint回调确实是“或”逻辑——每个回调各自独立判断自己的指标是否更优,只要满足就会触发保存,完全达不到“两个指标都要更优”的要求。要实现“与”逻辑,必须自定义一个回调函数,在每个epoch结束时同时检查两个指标的情况。
以下是适配Python 3.7和TensorFlow 2.1.0的实现方案:
自定义回调函数实现
import tensorflow as tf from tensorflow.keras.callbacks import Callback class DualMetricModelCheckpoint(Callback): def __init__(self, save_path, monitor_precision='val_precision', monitor_recall='val_recall'): super(DualMetricModelCheckpoint, self).__init__() self.save_path = save_path # 监控的指标名称,根据你的训练配置调整(比如不用val_前缀就是训练集指标) self.monitor_precision = monitor_precision self.monitor_recall = monitor_recall # 初始化最佳指标值,因为Precision和Recall范围是0-1,初始设为0即可 self.best_precision = 0.0 self.best_recall = 0.0 def on_epoch_end(self, epoch, logs=None): logs = logs or {} # 获取当前epoch的两个指标值 current_precision = logs.get(self.monitor_precision) current_recall = logs.get(self.monitor_recall) if current_precision is None or current_recall is None: raise ValueError(f"监控的指标 {self.monitor_precision} 或 {self.monitor_recall} 未在训练中输出,请检查模型编译时的metrics配置") # 判断是否两个指标都优于之前的最佳值 if current_precision > self.best_precision and current_recall > self.best_recall: print(f"\nEpoch {epoch+1}: {self.monitor_precision} 和 {self.monitor_recall} 均优于历史最佳,保存模型到 {self.save_path}") # 保存整个模型(也可以用save_weights只存权重,根据需求调整) self.model.save(self.save_path) # 更新最佳指标值 self.best_precision = current_precision self.best_recall = current_recall else: print(f"\nEpoch {epoch+1}: 未满足双指标均更优的条件,不保存模型")
使用方法
- 首先确保你的模型编译时已经加入了Precision和Recall指标:
# 注意TF2.1.0里的Precision/Recall类位置 from tensorflow.keras.metrics import Precision, Recall model.compile( optimizer='adam', loss='binary_crossentropy', # 根据你的任务调整损失函数 metrics=[Precision(name='precision'), Recall(name='recall')] )
- 在训练时加入自定义回调:
# 初始化自定义回调,指定保存路径 dual_checkpoint = DualMetricModelCheckpoint(save_path='best_dual_model.h5') # 开始训练 model.fit( train_data, validation_data=val_data, # 如果你用训练集指标,就不需要这个,同时回调里去掉val_前缀 epochs=50, callbacks=[dual_checkpoint] )
关键逻辑解释
- 自定义回调继承自
Callback,重写on_epoch_end方法,这个方法会在每个epoch训练结束后触发 - 初始化时记录最佳Precision和Recall的初始值(0.0)
- 每个epoch结束后,从
logs字典中获取当前的指标值(logs会包含训练时所有输出的指标) - 只有当当前Precision > 历史最佳Precision且当前Recall > 历史最佳Recall时,才保存模型并更新最佳值
- 如果其中一个指标没提升,或者两个都没提升,就不执行保存操作
内容的提问来源于stack exchange,提问作者Mistapopo
相关产品推荐
相关产品推荐

