如何用pickle保存Keras LSTM模型避免重训,遇弱引用无法序列化报错如何解决?
报错原因
基于TensorFlow后端的Keras模型包含大量与计算图绑定的弱引用、会话关联对象,这类对象本身不支持pickle序列化,因此直接用pickle转储Keras模型实例必然会触发cannot pickle 'weakref' object报错,普通Python弱引用的序列化方案不适用于该场景。
正确持久化Keras LSTM模型的方案
方案1:使用Keras原生H5格式保存(日常开发推荐)
直接用Keras内置接口保存完整模型,包含网络结构、权重、编译配置,无需依赖pickle。
- 保存代码:
# 模型训练完成后直接调用save方法 model.save('lstm_model.h5')
- 加载代码:
from tensorflow.keras.models import load_model loaded_model = load_model('lstm_model.h5') # 加载完成后可直接用于预测、增量训练
方案2:SavedModel格式保存(部署场景推荐)
如果后续需要对接TensorFlow生态的部署工具,使用官方推荐的SavedModel格式保存:
- 保存代码:
model.save('lstm_saved_model')
- 加载代码:
from tensorflow.keras.models import load_model loaded_model = load_model('lstm_saved_model')
特殊场景:必须用pickle打包模型和其他数据
如果业务逻辑要求必须将模型与其他元数据打包到字典中用pickle转储,可以先将模型序列化到内存字节流再存储:
import pickle from io import BytesIO # 序列化模型到内存字节流 model_byteio = BytesIO() model.save(model_byteio) model_byteio.seek(0) # 打包到字典后用pickle存储 mod = { 'Model': model_byteio.read(), # 可添加其他需要同步存储的元数据 'train_epochs': 1, 'batch_size': 1 } with open('lstm_model.pkl', 'wb') as file: pickle.dump(mod, file) # 加载时反向操作即可 with open('lstm_model.pkl', 'rb') as file: loaded_mod = pickle.load(file) loaded_model = load_model(BytesIO(loaded_mod['Model']))
内容的提问来源于stack exchange,提问作者nasc
相关产品推荐
相关产品推荐

