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

如何在tflite_model_maker模型训练时自动保存最优模型并设置早停

tflite_model_maker的目标检测模块本身就支持自动保存最优模型,也可以通过标准Keras回调配置早停,不需要手动监控损失中断训练,具体配置方式如下:

自动保存最优模型

object_detector.create() 默认已开启最优模型自动保存逻辑:训练过程中会持续追踪验证集损失,自动留存验证损失最低的模型检查点,全部训练轮次结束后,会自动加载该最优检查点作为返回的模型对象,无需额外编码。
如果需要自定义最优模型的判定规则、保存路径,可以通过HParams传入自定义配置:

from tflite_model_maker import object_detector
from tflite_model_maker.object_detector import HParams

custom_hparams = HParams(
    model_dir="./custom_checkpoint_dir",  # 检查点存储路径
    best_model_metric="val_map",  # 可选,默认判定指标为val_loss,可替换为验证集mAP等精度指标
    best_model_metric_mode="max"  # 指标最优判定方向:损失类指标用min,精度类指标用max
)

model = object_detector.create(
    train_data=train_dataset,
    validation_data=val_dataset,
    hparams=custom_hparams
)

配置早停触发

早停逻辑可以通过传入TensorFlow Keras原生的EarlyStopping回调实现,create()方法原生支持callbacks参数接收自定义回调列表,配置后训练过程中满足触发条件就会自动终止训练,不需要手动操作。
示例配置如下:

import tensorflow as tf
from tflite_model_maker import object_detector

# 定义早停规则
early_stop = tf.keras.callbacks.EarlyStopping(
    monitor="val_loss",  # 监控验证集损失
    patience=3,  # 连续3轮指标无提升就触发早停,可根据数据集规模调整
    restore_best_weights=True,  # 触发早停后自动加载最优轮次的权重,避免拿到过拟合的最后一轮模型
    verbose=1
)

# 训练时传入回调即可
model = object_detector.create(
    train_data=train_dataset,
    validation_data=val_dataset,
    epochs=100,  # 可以将总训练轮次设为较大值,早停触发后会自动终止
    callbacks=[early_stop]
)

小提示:如果你的数据集规模较小,建议把patience设为2-5,避免无效训练;如果数据集较大,可以适当调大到5-10,防止过早停止错过最优结果。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 09:18:19