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

训练Keras深度神经网络时,如何保存与加载keras.callbacks.History对象?

保存并加载Keras训练的History对象

没问题,要保存keras.callbacks.History类型的对象并在另一个Python会话中复用,最常用的是序列化工具(比如pickle或joblib),因为History对象的核心数据(训练/验证的loss、准确率等)都存在它的history属性里,这是一个可序列化的字典,完全满足后续分析需求。

下面是两种实用方法:

方法1:使用pickle(Python标准库)

保存History数据

训练完成后,直接把history属性序列化保存到文件:

import pickle

# 训练模型得到history对象
history_model_1 = model_1.fit_generator(train_generator, steps_per_epoch=100, epochs=20, validation_data=validation_generator, validation_steps=50)

# 保存history字典到文件
with open('history_model_1.pkl', 'wb') as file:
    pickle.dump(history_model_1.history, file)

加载并使用

在新的Python会话中,加载文件并恢复数据:

import pickle
import matplotlib.pyplot as plt

# 加载保存的history数据
with open('history_model_1.pkl', 'rb') as file:
    loaded_history = pickle.load(file)

# 示例:绘制训练/验证准确率曲线
plt.plot(loaded_history['accuracy'], label='Training Accuracy')
plt.plot(loaded_history['val_accuracy'], label='Validation Accuracy')
plt.xlabel('Epoch')
plt.ylabel('Accuracy')
plt.title('Model Accuracy Over Epochs')
plt.legend()
plt.show()

方法2:使用joblib(适合大数据量)

如果你的训练指标包含大量numpy数组,joblib比pickle更高效,它是scikit-learn附带的工具:

保存History数据

import joblib

# 训练后保存
joblib.dump(history_model_1.history, 'history_model_1.joblib')

加载并使用

import joblib
import matplotlib.pyplot as plt

loaded_history = joblib.load('history_model_1.joblib')

# 同样可以用来绘图或分析指标
plt.plot(loaded_history['loss'], label='Training Loss')
plt.plot(loaded_history['val_loss'], label='Validation Loss')
plt.xlabel('Epoch')
plt.ylabel('Loss')
plt.title('Model Loss Over Epochs')
plt.legend()
plt.show()

注意事项

  • 优先保存history_model_1.history(字典)而非整个History对象:History对象本身可能包含一些和当前训练会话绑定的内部状态,跨环境序列化容易出问题,而字典形式的指标数据是完全独立的,兼容性更好。
  • 确保保存和加载时使用相同的Keras/TensorFlow版本,避免因数据格式差异导致加载失败。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 07:52:19