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

TensorFlow使用saved_model.load加载模型预测报错解决

问题描述

我需要使用TensorFlow中已保存的模型完成预测,导出模型的实现代码如下:

def serving_input_fn_builder(max_seq_length):
    def _serving_input_fn():
        feature_spec = {
            "input_ids": tf.placeholder(tf.int32,
                                        shape=[None, max_seq_length]),
            "input_mask": tf.placeholder(tf.int32,
                                         shape=[None, max_seq_length]),
            "segment_ids": tf.placeholder(tf.int32,
                                          shape=[None, max_seq_length])
        }
        return tf.estimator.export.build_raw_serving_input_receiver_fn(
            feature_spec)()
    return _serving_input_fn

def export_tf_model(estimator, serving_input_fn, export_name="export"):
    estimator.export_savedmodel(
        export_dir_base=f"{export_name}",
        serving_input_receiver_fn=serving_input_fn)

调用导出逻辑的代码:

serving_input_fn = serving_input_fn_builder(FLAGS.max_seq_length)
export_model_params = {
    "export_name": "export"
}
export_tf_model(estimator, serving_input_fn, **export_model_params)

导出完成后得到saved_model.pb文件,尝试加载模型执行预测的代码如下:

checkpoint_path = "./Saved_Model_tf/saved_model.pb"
checkpoint_dir = os.path.dirname(checkpoint_path)

new_model = tf.saved_model.load(checkpoint_dir)
infer = new_model.signatures["serving_default"]
print(infer) ## 打印结果: {'prob': <tf.Tensor 'Softmax:0' shape=(None, 3) dtype=float32>, 'logits': <tf.Tensor 'loss/BiasAdd:0' shape=(None, 3) dtype=float32>, 'intent': <tf.Tensor 'ArgMax:0' shape=(None,) dtype=int32> => intent为预测结果

def predict(x):
    example = tf.train.Example()
    example.features.feature["sentence"].bytes_list.value.extend([x])
    out = new_model.signatures["predict"](examples=tf.constant([example.SerializeToString()]))['probabilities']
    return out

x = "The movie is awesome"
predict(x)

运行代码触发报错:

return self._signatures[key] KeyError: 'predict'

已查阅相关社区问答与官方指南资料仍未解决,需要正确的模型加载与推理实现方案。

报错原因
  • 模型导出时使用build_raw_serving_input_receiver_fn定义的输入签名为input_ids、input_mask、segment_ids三个int32类型张量,仅生成了默认的serving_default签名,不存在名为predict的签名,访问不存在的签名键直接触发KeyError。
  • 现有预测逻辑构造tf.train.Example序列化输入、取probabilities输出字段的写法,和模型实际的输入输出定义完全不匹配:该模型不接收序列化Example格式输入,输出字段中也不存在probabilities键。
  • 加载路径存在问题:tf.saved_model.load需要传入saved_model.pb所在的目录路径,且目录下必须包含导出时生成的variables子文件夹,直接指向pb文件本身可能导致权重加载异常。
正确实现代码
import tensorflow as tf
import numpy as np

# 替换为实际的模型目录:导出后export目录下会生成时间戳命名的子文件夹,该文件夹内包含saved_model.pb和variables目录
MODEL_DIR = "./export/你的模型时间戳文件夹"
MAX_SEQ_LENGTH = 你设置的max_seq_length值 # 和导出模型时传入的序列长度保持一致

# 加载模型
model = tf.saved_model.load(MODEL_DIR)
infer_func = model.signatures["serving_default"]

def predict(input_ids, input_mask, segment_ids):
    """
    三个入参均为shape=[batch_size, MAX_SEQ_LENGTH]的int32类型数组/张量
    """
    # 构造符合签名要求的输入张量
    input_tensors = {
        "input_ids": tf.constant(input_ids, dtype=tf.int32),
        "input_mask": tf.constant(input_mask, dtype=tf.int32),
        "segment_ids": tf.constant(segment_ids, dtype=tf.int32)
    }
    # 执行推理
    output = infer_func(**input_tensors)
    # 返回numpy格式的结果,按需取对应字段
    return {
        "prob": output["prob"].numpy(), # 分类概率
        "logits": output["logits"].numpy(), # 输出层原始值
        "intent": output["intent"].numpy() # 最终分类结果
    }

# 测试示例:替换为你实际分词逻辑生成的对应输入
# 单条输入shape为[1, MAX_SEQ_LENGTH]
test_input_ids = np.zeros((1, MAX_SEQ_LENGTH), dtype=np.int32)
test_input_mask = np.zeros((1, MAX_SEQ_LENGTH), dtype=np.int32)
test_segment_ids = np.zeros((1, MAX_SEQ_LENGTH), dtype=np.int32)

result = predict(test_input_ids, test_input_mask, test_segment_ids)
print(result)
注意事项
  • 传入tf.train.Example的调用方式仅适用于用build_parsing_serving_input_receiver_fn导出、接收序列化Example输入的模型,当前导出的是原始张量输入的模型,无需做Example序列化。
  • 推理时传入的参数名必须和导出时feature_spec定义的键名完全一致,输出取值的键名必须和打印serving_default时展示的输出键名完全一致。
  • 输入序列长度必须和导出模型时设置的max_seq_length保持一致,否则会触发shape不匹配错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.01 01:27:29