使用tflite_model_maker训练模型如何绘制准确率等性能变化曲线
解决方法
tflite_model_maker 封装了Keras训练逻辑,返回的模型对象默认不直接暴露history属性,你可以通过传入Keras原生回调的方式捕获训练过程的指标数据,具体操作如下:
- 提前导入依赖库并初始化历史记录回调
import matplotlib.pyplot as plt from tensorflow.keras.callbacks import History # 实例化History回调,用于捕获训练指标 history_cb = History()
- 修改
image_classifier.create调用,新增callbacks参数传入回调实例
model = image_classifier.create(train_data, model_spec = model_spec.get('efficientnet_lite4'), validation_data=validation_data, batch_size = 32, epochs=200, train_whole_model = True, dropout_rate=0.25, learning_rate = 0.01, momentum = 0.9, shuffle=True, # 新增回调参数,可同时传入多个Keras标准回调 callbacks=[history_cb] )
- 训练完成后,从回调对象中提取指标绘制曲线,示例代码如下:
plt.rcParams['font.sans-serif'] = ['SimHei'] # 可选,解决中文显示问题 plt.figure(figsize=(12,5)) # 绘制准确率曲线 plt.subplot(1,2,1) plt.plot(history_cb.history['accuracy'], label='训练准确率') plt.plot(history_cb.history['val_accuracy'], label='验证准确率') plt.title('准确率变化曲线') plt.xlabel('训练轮次') plt.ylabel('准确率') plt.legend() # 绘制损失曲线 plt.subplot(1,2,2) plt.plot(history_cb.history['loss'], label='训练损失') plt.plot(history_cb.history['val_loss'], label='验证损失') plt.title('损失变化曲线') plt.xlabel('训练轮次') plt.ylabel('损失值') plt.legend() plt.tight_layout() plt.show()
如果你需要早停、训练过程日志持久化等能力,也可以把EarlyStopping、TensorBoard等其他Keras标准回调一同添加到callbacks列表中,使用逻辑和原生Keras训练完全一致。
内容的提问来源于stack exchange,提问作者zeewa42
相关产品推荐
相关产品推荐

