TensorFlow Serving REST API传递Decoder模型多输入的正确方法
问题描述
参考Keras的seq2seq教程构建模型后,单独保存了encoder_model和decoder_model并通过TensorFlow Serving部署。编码器输入为长度1的numpy数组,转JSON调用REST API可正常运行,但解码器需要3个numpy数组输入时调用失败。
当前调用decoder_model的代码片段:
dec_model_url = "http://localhost:8400/v1/models/dec_model:predict" headers = { 'content-type': "application/json;charset=UTF-8'", 'cache-control': "no-cache", 'Accept':'application/json' } while not stop_condition: decoder_ip = ([target_seq] + states_value) target_seq1 = target_seq.tolist() target_seq1=[target_seq1] states_value1 = states_value states_value1[0] = states_value1[0].tolist() states_value1[1] = states_value1[1].tolist() decoder_ip1 = (target_seq1 + states_value1[0] + states_value1[1]) start_main = '{"instances":' end_main = '}' decoder_ip1 = start_main + str(decoder_ip1) +end_main output_tokens, h = requests.request("POST", dec_model_url, data=decoder_ip1, headers=headers)
运行时触发错误:
{"error": "instances is a plain list, but expecting list of objects as multiple input tensors required as per tensorinfo_map"}
解决方案
错误根源是TensorFlow Serving处理多输入模型时,要求instances是对象列表(每个对象对应一组输入张量的键值对),而非将所有输入扁平合并为一个列表。
正确做法是为每个输入张量指定对应名称(需与模型输入层的name属性一致),构造包含键值对的对象后放入instances数组中。
修正后的代码
import requests import json dec_model_url = "http://localhost:8400/v1/models/dec_model:predict" headers = { 'content-type': "application/json;charset=UTF-8", 'cache-control': "no-cache", 'Accept':'application/json' } while not stop_condition: # 将各输入转为列表格式 target_seq_list = target_seq.tolist() state_h_list = states_value[0].tolist() state_c_list = states_value[1].tolist() # 构造符合要求的请求体:instances为对象列表,每个对象包含所有输入张量的键值对 # 注意:键名必须与decoder_model输入层的name完全匹配 request_body = { "instances": [ { "target_input": target_seq_list, "state_h": state_h_list, "state_c": state_c_list } ] } # 使用json.dumps序列化,避免手动拼接JSON的语法错误 response = requests.post(dec_model_url, data=json.dumps(request_body), headers=headers) # 根据模型输出结构解析响应 result = response.json() output_tokens = result["predictions"][0][0] # 按需调整索引 h = result["predictions"][0][1] # 按需调整索引获取状态值
关键注意事项
- 输入名称匹配:请求体中的键名必须和
decoder_model输入层的name完全一致,可通过以下代码查看输入名称:print([input_layer.name for input_layer in decoder_model.inputs]) - 避免手动拼接JSON:使用
json.dumps()自动处理序列化,防止引号转义、格式不规范等问题。 - 响应解析适配:根据
decoder_model的输出结构,调整result["predictions"]的索引,确保正确提取output_tokens和状态值。
内容的提问来源于stack exchange,提问作者IronMan
相关产品推荐
相关产品推荐

