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
相关产品推荐
相关产品推荐

