无需训练代码恢复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
相关产品推荐
相关产品推荐

