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

如何用Keras的model.fit_generator()实现训练指标实时可视化?

解决方法:利用fit_generator返回的History对象实现可视化

嘿,其实你完全可以用和model.fit()几乎一样的方式实现训练指标可视化!model.fit_generator()同样会返回一个History对象,里面存着每一轮训练的损失、准确率(以及你指定的其他指标)和对应验证集数据,咱们只需要把这些数据提取出来画图就行。我来给你调整代码,一步步实现:

第一步:修改训练函数,返回History对象

先调整你的train_model函数,让它返回fit_generator的执行结果——这是实现可视化的核心前提:

from keras.callbacks import ModelCheckpoint

def train_model(model, args, X_train, X_valid, y_train, y_valid):
    # 假设你已经定义了自己的训练/验证数据生成器
    train_generator = ...  # 你的训练批次数据生成逻辑
    val_generator = ...    # 你的验证批次数据生成逻辑
    
    checkpoint = ModelCheckpoint(
        'model-{epoch:03d}.h5', 
        monitor='val_loss', 
        verbose=0, 
        save_best_only=args.save_best_only, 
        mode='auto', 
        period=1
    )
    
    # 执行训练并保存历史数据
    history = model.fit_generator(
        generator=train_generator,
        steps_per_epoch=len(X_train) // args.batch_size,
        validation_data=val_generator,
        validation_steps=len(X_valid) // args.batch_size,
        epochs=args.epochs,
        callbacks=[checkpoint]
    )
    
    return history  # 关键:返回训练历史对象

第二步:编写可视化函数

接下来写一个专门的绘图函数,把History里的指标数据转换成直观的曲线:

import matplotlib.pyplot as plt

def plot_training_metrics(history):
    # 设置画布大小,让图表更清晰
    plt.figure(figsize=(14, 6))
    
    # 绘制损失曲线
    plt.subplot(1, 2, 1)
    plt.plot(history.history['loss'], 'b-', label='训练损失')
    plt.plot(history.history['val_loss'], 'r--', label='验证损失')
    plt.title('训练与验证损失变化')
    plt.xlabel('轮次(Epoch)')
    plt.ylabel('损失值')
    plt.legend()
    plt.grid(True)
    
    # 绘制准确率曲线(如果模型编译时指定了accuracy指标)
    if 'accuracy' in history.history:
        plt.subplot(1, 2, 2)
        plt.plot(history.history['accuracy'], 'b-', label='训练准确率')
        plt.plot(history.history['val_accuracy'], 'r--', label='验证准确率')
        plt.title('训练与验证准确率变化')
        plt.xlabel('轮次(Epoch)')
        plt.ylabel('准确率')
        plt.legend()
        plt.grid(True)
    
    plt.tight_layout()
    plt.show()

第三步:调用训练与可视化

最后,只需要在训练完成后调用可视化函数,就能看到实时生成的指标曲线了:

# 启动训练并获取历史数据
training_history = train_model(model, args, X_train, X_valid, y_train, y_valid)

# 展示训练指标曲线
plot_training_metrics(training_history)

额外小提示

  • 如果你的模型编译时用的是acc而非accuracy(旧版Keras的习惯写法),记得把代码里的accuracy替换成acc,val_accuracy替换成val_acc。
  • 要是你想在每一轮训练结束后实时更新曲线(而非训练完再统一绘制),可以自定义一个Callback类,在on_epoch_end方法里更新绘图逻辑,不过一般训练完成后再绘制已经足够清晰直观。
  • 注意:在较新的Keras版本中,fit_generator已经被整合到model.fit()中(直接传入生成器即可使用),但你当前的用法依然完全兼容。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 04:35:40