如何在TensorFlow Lite Model Maker训练中定期保存检查点
用TensorFlow Lite Model Maker实现检查点保存与TensorBoard监控
要实现训练过程中定期保存检查点并通过TensorBoard查看损失曲线,你可以利用TensorFlow的Keras回调函数来搞定,因为object_detector.create其实支持传入callbacks参数,具体操作如下:
1. 导入必要的模块
先导入TensorFlow提供的两个核心回调类:
from tensorflow.keras.callbacks import ModelCheckpoint, TensorBoard import os
2. 创建保存目录
提前创建检查点和TensorBoard日志的保存目录,避免训练时出错:
# 自定义检查点和日志的保存路径 checkpoint_save_path = "./training_checkpoints" tensorboard_log_path = "./training_logs" # 自动创建目录(如果不存在) os.makedirs(checkpoint_save_path, exist_ok=True) os.makedirs(tensorboard_log_path, exist_ok=True)
3. 配置检查点回调
这个回调会帮你定期保存模型,你可以自定义保存频率、是否只存最优模型等:
checkpoint_callback = ModelCheckpoint( filepath=os.path.join(checkpoint_save_path, "detector_epoch_{epoch:02d}.h5"), save_weights_only=False, # False保存整个模型,True只保存权重,按需选择 save_freq="epoch", # 每训练1个epoch保存一次 monitor="val_loss", # 监控验证集损失,用来判断最优模型 save_best_only=True, # 只保存验证集损失最低的模型 verbose=1 # 保存时打印提示信息 )
4. 配置TensorBoard回调
这个回调会记录训练过程中的各项指标,方便后续在TensorBoard中查看:
tensorboard_callback = TensorBoard( log_dir=tensorboard_log_path, histogram_freq=1, # 每1个epoch记录一次权重直方图 write_graph=True, # 保存模型计算图 write_images=True, # 保存模型输入输出的可视化 update_freq="epoch" # 每epoch更新一次日志 )
5. 在训练时传入回调
修改你原本的object_detector.create调用,加入callbacks参数:
model = object_detector.create( train_data, model_spec=spec, batch_size=8, train_whole_model=True, validation_data=validation_data, callbacks=[checkpoint_callback, tensorboard_callback] )
6. 查看TensorBoard监控
训练过程中或训练结束后,在终端运行以下命令启动TensorBoard:
tensorboard --logdir=./training_logs
然后按照终端提示的地址(一般是http://localhost:6006)打开浏览器,就能看到损失曲线、模型结构等可视化内容了。
内容的提问来源于stack exchange,提问作者finisinfinitatis
相关产品推荐
相关产品推荐

