如何序列化与反序列化含Python对象的自定义TensorFlow模型
解决自定义TensorFlow模型序列化问题(兼容Python对象+TensorFlow张量)
方法1:用TensorFlow原生SavedModel格式(推荐)
tf.Module本身支持TensorFlow原生的SavedModel序列化,完美适配内部张量、层结构,不会碰到线程锁序列化的问题,是最稳妥的方案。
保存代码
import tensorflow as tf # 保存模型到指定目录 tf.saved_model.save(main_module, "./saved_main_module")
加载代码
# 从目录加载模型 loaded_module = tf.saved_model.load("./saved_main_module") # 验证功能正常 test_input = tf.random.normal((3, 2)) print(loaded_module(test_input))
优势:原生支持所有TensorFlow组件,无兼容性问题,模型还能跨语言加载。
方法2:自定义pickle序列化逻辑
如果一定要用pickle类库,需要手动跳过无法序列化的_thread.RLock对象(这类对象来自Keras层的线程安全机制),通过实现__reduce__和__setstate__方法,只保存必要的初始化参数和权重。
修改你的自定义Module类:
import tensorflow as tf import pickle class MyModule(tf.Module): def __init__(self, input_dim, output_dim): self.input_shape = input_dim self.dense = tf.keras.layers.Dense(output_dim, input_shape=(input_dim,)) # 原有属性、方法保持不变... def __reduce__(self): # 保存初始化所需参数+层的权重 init_args = (self.input_shape, self.dense.units) weights = self.dense.get_weights() # 返回格式:(构造函数, 构造参数, 待恢复的状态) return (MyModule, init_args, weights) def __setstate__(self, state): # 加载时恢复层的权重 self.dense.set_weights(state) class MyModuleWithSubModule(tf.Module): def __init__(self): self.sub_module = MyModule(input_dim=2, output_dim=1) # 原有属性、方法保持不变... def __reduce__(self): # 序列化子模块的状态 sub_module_state = pickle.dumps(self.sub_module) return (MyModuleWithSubModule, (), sub_module_state) def __setstate__(self, state): # 加载时恢复子模块 self.sub_module = pickle.loads(state)
序列化/反序列化测试
# 序列化模型为字节流 pickle_bytes = pickle.dumps(main_module) # 从字节流加载模型 loaded_module = pickle.loads(pickle_bytes) # 验证功能 test_input = tf.random.normal((3, 2)) print(loaded_module(test_input))
注意:需要给每个自定义tf.Module类都实现这两个方法,确保只序列化必要的参数和权重,跳过内部锁对象。
方法3:适配Keras模型保存
如果你的模块可以适配Keras结构,也可以用Keras的保存方法自动处理序列化:
# 将自定义模块包装成Keras模型 keras_model = tf.keras.Model( inputs=tf.keras.Input(shape=(2,)), outputs=main_module(tf.keras.Input(shape=(2,))) ) # 保存模型 tf.keras.models.save_model(keras_model, "./keras_main_module.h5") # 加载模型 loaded_keras_model = tf.keras.models.load_model("./keras_main_module.h5")
优势:无需手动写序列化逻辑,自动处理权重和层的保存,适合偏Keras风格的模型。
内容的提问来源于stack exchange,提问作者Stalin Thomas
相关产品推荐
相关产品推荐

