TensorFlow Object Detection API中如何存储最优模型检查点而非仅最新5个
嘿,这个问题我之前训练目标检测模型时也碰到过!确实挺闹心的——熬了一夜训练,结果最优模型被自动删了,只能看着过拟合的模型叹气。下面给你几个实用的解决方案:
解决TensorFlow Object Detection API检查点保存限制的方案
1. 直接增加保留的检查点数量
这是最快速的解决办法,修改你的训练配置文件(.config)里的train_config部分,调整keep_checkpoint_max参数:
train_config: { # 其他训练配置... keep_checkpoint_max: 20 # 把默认的5改成你想要保留的数量,比如20 # 其他训练配置... }
修改后重新启动训练,API就会保留最新的20个检查点,不用担心一夜训练后最优模型被清理掉。
2. 保存基于mAP的最优模型
如果想精准保存验证集mAP最高的模型,而不是只保留最新的,这里有两种靠谱的方法:
方法A:用API自带的BestCheckpointExporter(推荐)
TensorFlow Object Detection API已经内置了这个功能,你只需要两步配置:
- 在你的
.config文件的eval_config部分,确保启用了检测指标:
eval_config: { metrics_set: "coco_detection_metrics" # 针对检测任务的标准mAP指标 use_moving_averages: false }
- 启动评估进程时,加上
--export_best_exporter参数(比如用model_main_tf2.py的话):
python model_main_tf2.py --model_dir=你的训练目录 --pipeline_config_path=你的配置文件 --export_best_exporter
这样评估进程会自动跟踪验证集的mAP指标,把最优的检查点保存到你的训练目录/exported_models/best_checkpoint目录下,完全不用手动干预。
方法B:自定义回调函数(适合灵活需求)
如果你用的是TF2版本,也可以自己写一个回调函数来监控mAP,当指标提升时保存模型。示例代码如下:
import tensorflow as tf from object_detection.utils import metrics from object_detection.builders import dataset_builder from object_detection.utils import config_util class BestMapCallback(tf.keras.callbacks.Callback): def __init__(self, config_path, save_dir): super().__init__() # 加载配置与验证数据集 configs = config_util.get_configs_from_pipeline_file(config_path) eval_config = configs['eval_config'] eval_input_config = configs['eval_input_config'] self.val_dataset = dataset_builder.build( eval_input_config, eval_config, configs['model']) self.save_dir = save_dir self.best_map = 0.0 tf.io.gfile.makedirs(save_dir) def on_epoch_end(self, epoch, logs=None): # 计算当前验证集mAP map_metric = metrics.MeanAveragePrecision() for batch in self.val_dataset: images, gt_labels = batch predictions = self.model.predict(images) map_metric.update_state(gt_labels, predictions) current_map = map_metric.result().numpy() # 如果当前mAP更高,保存模型 if current_map > self.best_map: self.best_map = current_map print(f"🎉 新的最优mAP: {self.best_map:.4f}, 正在保存模型到 {self.save_dir}") checkpoint_path = f"{self.save_dir}/best_model_epoch_{epoch}" self.model.save_weights(checkpoint_path)
训练时把这个回调加入到训练流程里就行,能完全自定义保存逻辑。
3. 额外保险:手动定时备份
如果怕配置出问题,还可以写个简单的shell脚本,定时备份训练目录下的检查点文件(包括.index、.data-*、.meta这些),比如每10个epoch备份一次,这样就算默认清理了旧检查点,你也有备份可以恢复。
内容的提问来源于stack exchange,提问作者Piotr Januszewski
相关产品推荐
相关产品推荐

