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

无需训练代码恢复TensorFlow Eager模式模型的方法

在TensorFlow Eager模式下无需实例化Model类保存/加载模型的方案

问题回顾

你在Eager模式下训练了一个自定义tf.keras.Model,用Checkpoint保存后发现加载必须依赖原Model类的定义,而model.save()又抛出NotImplementedError,想要找到不依赖模型类实例化的保存加载方法。

可行解决方案:使用SavedModel格式

SavedModel是TensorFlow的标准模型序列化格式,它会完整保存模型的结构、权重和计算逻辑,加载时无需提前定义原模型类。针对你使用的TensorFlow 1.x Eager环境,具体操作如下:

1. 训练后导出为SavedModel

在你的训练代码末尾,添加以下代码来导出模型:

if __name__ == "__main__":
    # ... 保留原有的训练和Checkpoint保存代码 ...
    
    # 定义服务输入函数,指定模型输入的格式
    def serving_input_receiver_fn():
        # 定义输入占位符,匹配你的模型输入形状(这里是一维张量)
        input_ph = tf.placeholder(tf.float32, shape=[None], name="input")
        # 返回接收输入的对象
        return tf.contrib.eager.SavedModelInput(input_ph)
    
    # 导出模型到指定目录
    saved_model_dir = "./my_saved_model"
    tfe.save_saved_model(
        obj=model,
        export_dir=saved_model_dir,
        serving_input_fn=serving_input_receiver_fn,
        # 定义预测签名,明确输入输出映射
        signature_def_map={
            tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY:
                tfe.signature_def_utils.predict_signature_def(
                    inputs={"input": input_ph},
                    outputs={"output": model.predict(input_ph)}
                )
        }
    )

2. 无需原Model类即可加载并推理

在另一个脚本中,直接加载SavedModel并进行预测,完全不需要定义原Model类:

import tensorflow as tf
import tensorflow.contrib.eager as tfe

tf.enable_eager_execution()

# 加载SavedModel
loaded_model = tfe.load_saved_model("./my_saved_model")

# 获取默认的预测签名函数
predict_signature = loaded_model.signatures[
    tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY
]

# 执行预测,输入为Tensor类型
test_input = tf.constant(7., dtype=tf.float32)
prediction = predict_signature(input=test_input)
print(f"预测结果:{prediction['output'].numpy()}")

针对TensorFlow 2.x的补充

如果你升级到TF2.x(Eager模式默认启用),API会更简洁,直接使用原生的tf.saved_model模块即可:

# TF2.x保存模型
tf.saved_model.save(model, "./tf2_saved_model")

# TF2.x加载并预测
loaded_model = tf.saved_model.load("./tf2_saved_model")
result = loaded_model(tf.constant(7.))
print(result.numpy())

为什么这个方案可行?

SavedModel不像Checkpoint只保存变量权重,它会把模型的计算图、输入输出规范、权重数据全部序列化存储,因此加载时不需要依赖原模型类的定义,直接就能还原出可执行的模型计算逻辑。

内容的提问来源于stack exchange,提问作者Thomas Reynaud

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.27 09:30:24