使用pickle保存Keras ANN回归模型报can't pickle weakref objects错误
错误原因
Keras(TensorFlow后端)的Sequential模型对象内部包含弱引用(weakref)类型的属性,Python标准库pickle不支持序列化这类对象,直接调用pickle.dump()存储整个模型就会触发TypeError: can't pickle weakref objects报错。
解决方法
不要直接用原生pickle序列化整个Keras模型对象,根据使用场景选以下任意一种方案即可:
方案1:使用Keras原生保存接口(官方推荐)
Keras自带的保存接口可以完整存储模型结构、权重、训练配置,不存在序列化兼容问题,是优先选择的方案。
将原代码中pickle相关的保存、加载逻辑替换为以下代码即可:
# 先在代码头部导入加载模型的方法 from tensorflow.keras.models import load_model # 模型训练完成后,替换原pickle保存逻辑 modell.save('model.h5') # 加载模型 model = load_model('model.h5')
用这种方式加载的模型保留了全部训练配置,不需要重新compile,可以直接用于预测、评估或者继续训练。
方案2:必须使用pickle格式的兼容写法
如果业务场景要求必须输出pkl格式文件,可以将模型结构和权重拆分后再用pickle存储,不要直接序列化整个模型对象:
import pickle # 保存逻辑:拆分模型结构和权重后存储 model_config = modell.to_json() model_weights = modell.get_weights() pickle.dump( {'config': model_config, 'weights': model_weights}, open('model.pkl', 'wb') ) # 加载逻辑:读取后重构模型 from tensorflow.keras.models import model_from_json saved_data = pickle.load(open('model.pkl', 'rb')) model = model_from_json(saved_data['config']) model.set_weights(saved_data['weights']) # 注意:该方式加载的模型需要重新compile才能用于评估、继续训练 model.compile(loss='mean_squared_error', optimizer='adam', metrics=['mean_squared_error'])
内容的提问来源于stack exchange,提问作者CO19 301 Aanchal Bhatti
相关产品推荐
相关产品推荐

