如何可视化tflite-model-maker中image_classifier.create()的模型训练过程?
TFLite Model Maker训练过程可视化实现方案
报错原因说明
你执行model = model_spec.get('efficientnet_lite0')触发报错,是因为model_spec不是可直接全局调用的对象,需要先从对应模块导入才可以使用。同时这种手动实例化模型、调用compile和fit的写法完全脱离了TFLite Model Maker的封装逻辑,你手动训练的模型和用image_classifier.create训练的模型是两个完全独立的对象,自然也拿不到实际训练过程的历史数据。
训练后可视化方案
image_classifier.create方法返回的模型实例本身自带Keras训练历史属性,不需要修改原有训练逻辑,训练结束后直接提取model.history即可生成可视化曲线,示例代码如下:
import matplotlib.pyplot as plt # 直接提取训练历史 history = model.history # 绘制准确率曲线 plt.figure(figsize=(12, 4)) plt.subplot(1, 2, 1) plt.plot(history.history['accuracy'], label='训练准确率') plt.plot(history.history['val_accuracy'], label='验证准确率') plt.title('模型准确率') plt.xlabel('训练轮次') plt.ylabel('准确率') plt.legend() # 绘制损失曲线 plt.subplot(1, 2, 2) plt.plot(history.history['loss'], label='训练损失') plt.plot(history.history['val_loss'], label='验证损失') plt.title('模型损失') plt.xlabel('训练轮次') plt.ylabel('损失值') plt.legend() plt.tight_layout() plt.show()
训练过程实时可视化方案
如果需要在训练过程中实时查看指标变化,可以给create方法传入TensorBoard回调,操作步骤如下:
- 导入并配置TensorBoard回调
from tensorflow.keras.callbacks import TensorBoard import datetime # 配置日志保存路径 log_dir = "logs/fit/" + datetime.datetime.now().strftime("%Y%m%d-%H%M%S") tensorboard_callback = TensorBoard(log_dir=log_dir, histogram_freq=1)
- 把回调参数加到原有训练代码的
create方法中
model = tflite_model_maker.image_classifier.create( train_data, model_spec='efficientnet_lite0', use_augmentation=True, validation_data=validation_data, epochs=30, dropout_rate=0.3, learning_rate=0.0001, shuffle=True, # 新增回调参数 callbacks = [tensorboard_callback] )
- 训练启动后,在终端执行命令
tensorboard --logdir logs/fit,按照终端提示打开浏览器访问对应本地地址,即可实时查看训练指标曲线。
注意:你不需要手动实例化ModelSpec、调用compile和fit方法,Model Maker的create接口已经封装了数据预处理、增强、模型编译、训练的全流程,手动写fit逻辑反而会丢失你配置的增强、自动预处理等特性。
内容的提问来源于stack exchange,提问作者ahmdksyrf
相关产品推荐
相关产品推荐

