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

如何可视化MediaPipe图像分类器的损失与准确率?

MediaPipe ImageClassifier训练指标可视化方案

问题说明

使用MediaPipe的image_classifier.ImageClassifier.create()训练图像分类模型后,想要绘制训练/验证的loss、accuracy曲线,但发现模型实例没有history属性,报错:

AttributeError: 'ImageClassifier' object has no attribute 'history'

解决方法

MediaPipe Model Maker并未直接暴露训练历史属性,但可以通过自定义Keras回调函数在训练过程中记录每轮的指标数据,后续用这些数据生成可视化图表。

步骤1:定义自定义回调

创建一个回调类,在训练开始和每个epoch结束时记录loss、accuracy等指标:

import tensorflow as tf

class TrainingHistory(tf.keras.callbacks.Callback):
    def on_train_begin(self, logs=None):
        self.history = {'loss': [], 'accuracy': [], 'val_loss': [], 'val_accuracy': []}
    
    def on_epoch_end(self, epoch, logs=None):
        self.history['loss'].append(logs['loss'])
        self.history['accuracy'].append(logs['accuracy'])
        self.history['val_loss'].append(logs['val_loss'])
        self.history['val_accuracy'].append(logs['val_accuracy'])

history_callback = TrainingHistory()

步骤2:添加回调到训练选项

将自定义回调传入训练配置的callbacks参数中:

from mediapipe.model_maker import image_classifier

# 配置训练参数(根据需求调整epoch数等)
options = image_classifier.ImageClassifierOptions(
    hparams=image_classifier.HParams(epochs=10),
    callbacks=[history_callback]
)

# 启动训练
model = image_classifier.ImageClassifier.create(
    train_data=train_data,
    validation_data=validation_data,
    options=options
)

步骤3:绘制指标曲线

利用回调记录的数据绘制loss和accuracy图表:

import matplotlib.pyplot as plt
%matplotlib inline

history_dict = history_callback.history
epochs = range(1, len(history_dict['loss']) + 1)

# 绘制Loss曲线
plt.figure(figsize=(10, 5))
plt.plot(epochs, history_dict['loss'], 'b+', linewidth=2, markersize=10, label='训练Loss')
plt.plot(epochs, history_dict['val_loss'], 'r*', linewidth=2, markersize=10, label='验证Loss')
plt.xlabel('Epochs')
plt.ylabel('Loss')
plt.grid(True)
plt.legend()
plt.show()

# 绘制Accuracy曲线
plt.figure(figsize=(10, 5))
plt.plot(epochs, history_dict['accuracy'], 'b+', linewidth=2, markersize=10, label='训练Accuracy')
plt.plot(epochs, history_dict['val_accuracy'], 'r*', linewidth=2, markersize=10, label='验证Accuracy')
plt.xlabel('Epochs')
plt.ylabel('Accuracy')
plt.grid(True)
plt.legend()
plt.show()

内容的提问来源于stack exchange,提问作者uh31234

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 23:55:16