Django项目pickle序列化Keras模型报can't pickle weakref objects错误如何解决
Django项目中pickle序列化Keras模型报错的解决方案
报错原因
Keras/TensorFlow模型对象内部包含弱引用类型的属性,Python原生pickle模块不支持序列化这类对象,所以执行pickle.dump存储模型时会抛出Type Error: can't pickle weakref objects错误。
解决方案
- 区分存储对象类型,Keras模型用TensorFlow自带的存储方法,不要用pickle序列化。scaler这类sklearn生成的对象仍可正常使用pickle存储。
- 修改模型存储路径的后缀,不要用
.pkl作为Keras模型的存储后缀,推荐使用.h5格式或者直接用SavedModel目录格式。 - 对应修改模型加载逻辑,加载Keras模型时使用
tf.keras.models.load_model()方法,scaler仍用pickle加载即可。
修改后的代码示例
utils.py
import pickle import tensorflow as tf # 将模型导出到指定路径 def store_model(path, model): # 判定如果是Keras模型,用官方接口存储 if isinstance(model, tf.keras.Model): model.save(path) return # 其他可序列化对象继续用pickle存储 with open(path, 'wb') as f: pickle.dump(model, f)
build.py
# 存储模型部分修改为以下代码 store_model(f'{path}/model.h5', model) store_model(f'{path}/scaler_in.pkl', scaler_in) store_model(f'{path}/scaler_out.pkl', scaler_out)
注意事项
优先使用TensorFlow官方提供的存储、加载接口,可避免绝大多数序列化相关问题,同时兼容性更好,适配不同版本的TensorFlow环境。存储模型的操作要放在调用keras_clear()之前执行,避免模型被清理后存储空对象,现有代码的执行顺序无需调整。
内容的提问来源于stack exchange,提问作者Muhammad Fayzan
相关产品推荐
相关产品推荐

