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

如何在Keras中每个epoch保存训练历史以支持续训及图表生成?

当然可以!完全可以通过自定义Keras Callback来实现每个epoch保存训练历史,而且这正是解决你这种中断后继续训练、生成完整训练曲线需求的标准方案。我给你一步步拆解怎么做:

1. 自定义Callback实现每个Epoch保存训练历史

我们可以写一个继承自keras.callbacks.Callback的类,在每个epoch结束后把当前的训练指标(loss、accuracy等)追加保存到文件中。这里推荐用JSON格式,它易读且方便后续加载合并:

import json
from keras.callbacks import Callback

class SaveHistoryCallback(Callback):
    def __init__(self, filepath):
        super().__init__()
        self.filepath = filepath
        # 初始化时尝试加载已有的历史(如果之前保存过)
        try:
            with open(self.filepath, 'r') as f:
                self.history = json.load(f)
        except FileNotFoundError:
            # 根据你的训练指标调整这里的键名,比如用categorical_accuracy等
            self.history = {'loss': [], 'accuracy': [], 'val_loss': [], 'val_accuracy': []}

    def on_epoch_end(self, epoch, logs=None):
        logs = logs or {}
        # 将当前epoch的指标追加到历史字典中
        for metric_name, value in logs.items():
            if metric_name in self.history:
                self.history[metric_name].append(float(value))
        # 把更新后的历史写入文件
        with open(self.filepath, 'w') as f:
            json.dump(self.history, f)

这个Callback会自动处理历史的加载和追加,就算中途断电或关闭程序,下次启动时也能从上次的历史继续记录。

2. 搭配ModelCheckpoint保存模型权重

只保存历史还不够,你需要保存模型的训练状态,这样下次才能从上次的进度继续训练。用Keras自带的ModelCheckpoint就能轻松实现:

from keras.callbacks import ModelCheckpoint

# 每个epoch都更新保存模型权重(也可以设置只保存最优权重,看你的需求)
checkpoint_callback = ModelCheckpoint(
    filepath='model_weights.h5',
    save_weights_only=True,
    save_freq='epoch',
    verbose=1
)
3. 第一次训练(比如100个epoch)

训练时把两个Callback都加入到回调列表里即可:

# 假设你已经定义好了model、train_generator和val_generator
history_callback = SaveHistoryCallback(filepath='training_history.json')

model.fit_generator(
    generator=train_generator,
    epochs=100,
    validation_data=val_generator,
    callbacks=[history_callback, checkpoint_callback]
)
4. 次日继续训练(再50个epoch)

重启程序后,先加载之前保存的模型权重和训练历史,然后设置initial_epoch参数从上次结束的位置继续:

# 加载之前保存的模型权重
model.load_weights('model_weights.h5')

# 初始化历史Callback,它会自动加载已有的training_history.json
history_callback = SaveHistoryCallback(filepath='training_history.json')

# 继续训练,总epoch数设为150,初始epoch设为100(上次训练的结束点)
model.fit_generator(
    generator=train_generator,
    epochs=150,
    initial_epoch=100,
    validation_data=val_generator,
    callbacks=[history_callback, checkpoint_callback]
)
5. 生成完整的训练曲线

训练完成后,加载合并后的历史文件,用Matplotlib绘制完整的Loss和Accuracy曲线:

import matplotlib.pyplot as plt
import json

# 加载完整的训练历史
with open('training_history.json', 'r') as f:
    full_history = json.load(f)

epochs = range(1, len(full_history['loss']) + 1)

# 绘制Loss曲线
plt.figure(figsize=(12, 5))
plt.subplot(1, 2, 1)
plt.plot(epochs, full_history['loss'], 'b-', label='Training Loss')
plt.plot(epochs, full_history['val_loss'], 'r-', label='Validation Loss')
plt.title('Loss vs Epoch')
plt.xlabel('Epoch')
plt.ylabel('Loss')
plt.legend()

# 绘制Accuracy曲线
plt.subplot(1, 2, 2)
plt.plot(epochs, full_history['accuracy'], 'b-', label='Training Accuracy')
plt.plot(epochs, full_history['val_accuracy'], 'r-', label='Validation Accuracy')
plt.title('Accuracy vs Epoch')
plt.xlabel('Epoch')
plt.ylabel('Accuracy')
plt.legend()

plt.tight_layout()
plt.show()

小提示

  • 注意你的训练指标名称要和Callback里定义的键一致,比如如果用多分类任务,可能指标是categorical_accuracy,要对应修改。
  • 如果你用的是较新版本的Keras,fit_generator已经被fit替代(fit现在支持generator输入),直接替换成model.fit(...)即可,用法完全一致。

内容的提问来源于stack exchange,提问作者Preetom Saha Arko

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 10:15:38