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

TensorFlow ServingInputReceiver客户端输入形状错误及Estimator服务困惑

解决TensorFlow Estimator + TensorFlow Serving 接收序列化Proto的问题

我之前折腾Estimator和TensorFlow Serving的时候也被这个问题卡过,官方文档确实太敷衍了,完全没讲清楚端到端的流程。结合我踩过的坑,给你一步步拆解怎么实现:

1. 先定义你的输入Proto结构

首先你得有一个明确的.proto文件,用来定义要发送的序列化消息格式。比如假设你的输入包含用户ID和特征向量,创建一个input.proto:

syntax = "proto3";

message InputMessage {
  int32 user_id = 1;
  repeated float features = 2;
}

然后用protoc生成Python对应的代码:

protoc --python_out=. input.proto

这样你就得到input_pb2.py,用来在客户端序列化消息。

2. 在ServingInputReceiver中处理Proto字符串解析

Estimator导出模型时,ServingInputReceiver需要接收字符串类型的输入张量,然后在TensorFlow图里完成反序列化。这里关键是用tf.io.decode_proto函数直接在图中解析proto字符串,不用在客户端提前解析(这正是你要的效果)。

示例代码大概是这样:

import tensorflow as tf
from your_module import input_pb2  # 导入生成的proto代码

def serving_input_receiver_fn():
    # 定义接收的输入张量:形状是[None]的字符串张量,对应批量的序列化proto消息
    serialized_tf_example = tf.placeholder(
        dtype=tf.string,
        shape=[None],
        name='input'  # 这个name要和客户端发送时的输入键对应
    )
    
    # 定义proto解析的字段映射,要和你的.proto文件完全匹配
    field_names = ['user_id', 'features']
    output_types = [tf.int32, tf.float32]
    
    # 解析序列化的proto字符串
    parsed_result = tf.io.decode_proto(
        serialized_tf_example,
        message_type=input_pb2.InputMessage.DESCRIPTOR.full_name,
        field_names=field_names,
        output_types=output_types,
        descriptor_pool=tf.compat.v1.gfile.GFile(input_pb2.DESCRIPTOR.file.name).read()
    )
    
    # 把解析后的结果整理成模型需要的输入格式(比如VarLenFeature转成dense张量)
    inputs = {
        'user_id': parsed_result.values[0],
        'features': tf.sparse.to_dense(parsed_result.values[1])
    }
    
    return tf.estimator.export.ServingInputReceiver(inputs, {'input': serialized_tf_example})

# 导出模型
estimator.export_saved_model(
    export_dir_base='./saved_model',
    serving_input_receiver_fn=serving_input_receiver_fn
)

这里要注意几个点:

  • message_type必须是你的proto消息的完整名称(比如input.InputMessage,取决于你的proto包名)
  • descriptor_pool要正确加载你的proto文件描述符,否则TensorFlow找不到对应的消息定义
  • 解析后的特征要转换成模型输入需要的格式,比如稀疏张量转稠密张量

3. 用prediction_service_pb2发送序列化请求

客户端这边,你需要创建PredictRequest,把序列化的InputMessage放到inputs['input']的string_val里,然后通过gRPC调用TensorFlow Serving服务。

示例代码:

import grpc
from tensorflow_serving.apis import prediction_service_pb2_grpc
from tensorflow_serving.apis import predict_pb2
from your_module import input_pb2

# 连接TensorFlow Serving服务
channel = grpc.insecure_channel('localhost:8500')
stub = prediction_service_pb2_grpc.PredictionServiceStub(channel)

# 创建并序列化你的输入消息
input_msg = input_pb2.InputMessage()
input_msg.user_id = 123
input_msg.features.extend([0.1, 0.2, 0.3])
serialized_msg = input_msg.SerializeToString()

# 构建PredictRequest
request = predict_pb2.PredictRequest()
request.model_spec.name = 'your_model_name'  # 对应Serving启动时的模型名称
request.model_spec.signature_name = 'serving_default'  # 默认签名,或者你自定义的签名

# 把序列化消息放到输入张量里
request.inputs['input'].string_val.append(serialized_msg)

# 发送请求并获取响应
response = stub.Predict(request, timeout=10.0)

# 处理响应,比如获取预测结果
predictions = response.outputs['predictions'].float_val
print(predictions)

关键注意事项

  • 导出模型时,确保ServingInputReceiver的输入字典里的键(这里是'input')和客户端发送的request.inputs的键完全一致
  • 如果你自定义了签名名称,要在request.model_spec.signature_name里指定,否则用默认的serving_default
  • 测试时可以先用saved_model_cli验证模型的输入输出格式是否正确,比如:
saved_model_cli show --dir ./saved_model/1234567890 --all

查看输入张量的类型和名称是否符合预期

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.26 10:54:04