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

使用tflite_model_maker训练模型如何绘制准确率等性能变化曲线

解决方法

tflite_model_maker 封装了Keras训练逻辑,返回的模型对象默认不直接暴露history属性,你可以通过传入Keras原生回调的方式捕获训练过程的指标数据,具体操作如下:

  1. 提前导入依赖库并初始化历史记录回调
import matplotlib.pyplot as plt
from tensorflow.keras.callbacks import History

# 实例化History回调,用于捕获训练指标
history_cb = History()
  1. 修改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]
                                ) 
  1. 训练完成后,从回调对象中提取指标绘制曲线,示例代码如下:
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 22:45:07