You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何序列化与反序列化含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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.02 05:10:23