如何在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
相关产品推荐
相关产品推荐

