如何在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
相关产品推荐
相关产品推荐

