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

如何使用TensorFlow Object Detection API实现类Keras的交叉验证与留出验证并保存最优模型

实现方案

一、自动保存最优验证模型方案

你当前的并行训练+评估流程已经实现了留出验证的基础逻辑,只需要新增一个指标监听和最优检查点备份的脚本即可实现类似Keras的ModelCheckpoint功能,不需要修改官方源码:

步骤1:调整pipeline配置

在原有配置基础上新增以下参数,避免最优检查点被覆盖、保证评估流程正常运行:

train_config {
  # 原有训练配置保留,新增以下参数
  keep_checkpoint_max: 20 # 可自定义保留最近的检查点数量,避免最优检查点被删除
}

eval_config {
  metrics_set: "coco_detection_metrics"
  use_moving_averages: false
  batch_size: 1
  # 原有评估配置保留,新增以下参数
  max_evals: 0 # 不限制评估次数
  eval_interval_secs: 60 # 每隔60秒扫描一次是否有新检查点可评估
}

步骤2:新增最优模型监听脚本

在Colab中新增一个单元格,运行以下独立脚本,和训练、评估脚本并行执行即可自动备份最优模型:

import tensorflow as tf
import os
import shutil
import time

# 自定义配置项
MODEL_DIR = "替换为你的模型输出目录路径"
BEST_MODEL_SAVE_DIR = "./best_model"
# 可替换为你需要的目标指标,比如用mAP@0.5的话值为"DetectionBoxes_Precision/mAP@.50IOU"
TARGET_EVAL_METRIC = "DetectionBoxes_Precision/mAP"

os.makedirs(BEST_MODEL_SAVE_DIR, exist_ok=True)
best_score = 0.0
best_step = 0
processed_event_files = set()

while True:
    # 遍历所有评估事件文件
    for file_name in os.listdir(MODEL_DIR):
        if file_name.startswith("events.out.tfevents") and "eval" in file_name and file_name not in processed_event_files:
            processed_event_files.add(file_name)
            # 读取事件中的评估指标
            event_path = os.path.join(MODEL_DIR, file_name)
            for raw_record in tf.data.TFRecordDataset(event_path):
                event = tf.compat.v1.Event.FromString(raw_record.numpy())
                for val in event.summary.value:
                    if val.tag == TARGET_EVAL_METRIC:
                        current_score = val.simple_value
                        current_step = event.step
                        print(f"Step {current_step} 指标{TARGET_EVAL_METRIC}为{current_score:.4f}")
                        # 比对更新最优模型
                        if current_score > best_score:
                            best_score = current_score
                            best_step = current_step
                            print(f"更新最优模型:step {best_step},指标值{best_score:.4f}")
                            # 复制最优检查点到指定目录
                            for suffix in [".index", ".data-00000-of-00001", ".meta"]:
                                src_ckpt = os.path.join(MODEL_DIR, f"ckpt-{best_step}{suffix}")
                                if os.path.exists(src_ckpt):
                                    shutil.copy(src_ckpt, BEST_MODEL_SAVE_DIR)
                            # 同步复制配置文件
                            shutil.copy(os.path.join(MODEL_DIR, "pipeline.config"), BEST_MODEL_SAVE_DIR)
    # 每隔30秒扫描一次新的评估结果
    time.sleep(30)

二、多折交叉验证实现方案

基于上述留出验证的逻辑,按以下步骤扩展即可实现K折交叉验证:

  • 第一步:将全量数据集拆分为K份,生成K组对应的训练集TFRecord和验证集TFRecord,比如K=5时生成train_fold0.tfrecord/val_fold0.tfrecord到train_fold4.tfrecord/val_fold4.tfrecord
  • 第二步:为每折数据单独编写pipeline配置文件,分别修改train_input_reader和eval_input_reader的输入路径,同时为每折单独指定模型输出目录
  • 第三步:编写循环执行脚本,依次运行每折的训练、评估、最优模型筛选流程,每折跑完后记录最优指标,最后计算所有折的指标平均值即为交叉验证结果
    如果Colab显存不足以并行跑多折任务,可挂载Google Drive存储中间结果,依次串行执行每折任务即可。

内容的提问来源于stack exchange,提问作者F.M.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 22:06:04