TensorFlow 2.6 Object Detection如何保存所有checkpoints并运行全量检查点评估
TensorFlow 2.6 Object Detection 训练配置问题解决方案
1. 保存所有训练生成的checkpoints
TensorFlow Object Detection API默认最多保留最近7个checkpoint,可通过修改配置文件或自定义训练脚本参数调整:
- 基于官方pipeline配置文件启动训练的场景,找到
train_config配置段,修改keep_checkpoint_max参数:train_config { ... keep_checkpoint_max: 0 # 设为*0*代表不限制保存数量,所有checkpoint都会永久保留 ... } - 自定义训练脚本的场景,修改
CheckpointManager初始化参数,将max_to_keep设为None:checkpoint_manager = tf.train.CheckpointManager(checkpoint, directory=ckpt_save_dir, max_to_keep=None)
2. 对所有已保存的checkpoints分别执行评估
同样可通过配置文件或自定义脚本实现全量checkpoint评估:
- 基于官方pipeline配置文件启动评估的场景,找到
eval_config配置段,添加/修改run_all_checkpoints参数:eval_config { ... run_all_checkpoints: true # 开启后会按生成顺序逐个加载所有已保存的checkpoint执行评估 eval_interval_secs: 300 # 可根据训练生成checkpoint的频率调整评估间隔,避免重复读取 ... } - 自定义评估脚本的场景,遍历checkpoint目录下的所有历史 checkpoint 路径循环评估即可,示例逻辑:
# 获取所有历史checkpoint路径 ckpt_state = tf.train.get_checkpoint_state(ckpt_dir) all_ckpt_paths = ckpt_state.all_model_checkpoint_paths if ckpt_state else [] for ckpt_path in all_ckpt_paths: model.load_weights(ckpt_path).expect_partial() # 执行自定义评估逻辑 eval_metrics = run_evaluate(model, eval_dataset) print(f"Checkpoint {ckpt_path} 评估结果:{eval_metrics}")
3. 同时运行训练和评估任务时OOM错误的解决方案
OOM本质是显存不足以同时支撑两个任务的开销,可通过以下方式解决:
- 开启显存动态分配:在训练、评估脚本的开头添加配置,避免TensorFlow一次性占用全部显存:
也可以手动限制单任务的显存上限,比如给评估任务分配2G显存:gpus = tf.config.list_physical_devices('GPU') if gpus: for gpu in gpus: tf.config.experimental.set_memory_growth(gpu, True)tf.config.experimental.set_virtual_device_configuration(gpus[0], [tf.config.experimental.VirtualDeviceConfiguration(memory_limit=2048)]) - 降低评估批次大小:修改
eval_config中的eval_batch_size参数,设为1或者2,评估不需要反向传播,小batch size不会明显影响评估速度,但能大幅降低显存占用。 - 错开任务运行时间:不需要实时监控训练指标的场景,可等训练全部完成后再批量执行所有checkpoint的评估,避免两个高显存消耗任务同时运行。
- 拆分硬件资源:如果有多张GPU,通过
CUDA_VISIBLE_DEVICES环境变量指定训练和评估任务运行在不同GPU上,例如:
训练启动命令:CUDA_VISIBLE_DEVICES=0 python model_main_tf2.py --model_dir=./model --pipeline_config_path=./pipeline.config
评估启动命令:CUDA_VISIBLE_DEVICES=1 python model_main_tf2.py --model_dir=./model --pipeline_config_path=./pipeline.config --eval_training_data=False
内容的提问来源于stack exchange,提问作者harshata
相关产品推荐
相关产品推荐

