pickle反序列化TensorFlow模型后调用predict方法报错求助
解决方案
问题根源
报错核心原因是现有实现中,直接覆写Sequential实例的__dict__,跳过了Keras模型类初始化阶段内置私有属性的生成流程,导致_distribution_strategy等运行必需的私有属性缺失。
修复方案
将序列化/反序列化逻辑从Keras模型层面迁移到你的自定义类层面,避免hack Keras内部方法,修改后可直接运行的完整代码如下:
import pickle import numpy as np import tempfile import tensorflow as tf class ContainsSequential: def __init__(self): self.other_field = "potato" # 初始化模型逻辑不变 self.model = tf.keras.models.Sequential() self.model.add(tf.keras.layers.Input(shape=(None, 3))) self.model.add(tf.keras.layers.LSTM(3, activation="relu", return_sequences=True)) self.model.add(tf.keras.layers.Dense(3, activation="linear")) def __getstate__(self): # 序列化时:先存模型的二进制内容,再存其他所有字段 state = self.__dict__.copy() with tempfile.NamedTemporaryFile(suffix=".hdf5", delete=False) as fd: tf.keras.models.save_model(self.model, fd.name, overwrite=True) state["model_bin"] = fd.read() # 删掉原model对象,只保留二进制串 del state["model"] return state def __setstate__(self, state): # 反序列化时:先加载模型,再恢复其他字段 self.__dict__ = state with tempfile.NamedTemporaryFile(suffix=".hdf5", delete=False) as fd: fd.write(state["model_bin"]) fd.flush() self.model = tf.keras.models.load_model(fd.name) # 删掉临时存储的模型二进制内容 del self.__dict__["model_bin"] # 主逻辑执行: tf.keras.backend.clear_session() file_name = 'pickle_file.pckl' instance = ContainsSequential() instance.model.predict(np.random.rand(3, 1, 3)) print(instance.other_field) with open(file_name, 'wb') as fid: pickle.dump(instance, fid) with open(file_name, 'rb') as fid: restored_instance = pickle.load(fid) print(restored_instance.other_field) restored_instance.model.predict(np.random.rand(3, 1, 3)) print('Done')
方案说明
- 完全满足自定义类灵活序列化的需求,后续新增字段、方法都不需要调整序列化逻辑,兼容所有业务字段的正常存取
- 无需修改Keras内部属性、不hack模型类的内置方法,兼容性更强
- 反序列化后的模型可直接正常调用
predict方法执行预测,符合无需再训练的业务要求
内容的提问来源于stack exchange,提问作者Mefitico
相关产品推荐
相关产品推荐

