训练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
相关产品推荐
相关产品推荐

