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

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]  # 按需调整索引获取状态值

关键注意事项

  1. 输入名称匹配:请求体中的键名必须和decoder_model输入层的name完全一致,可通过以下代码查看输入名称:
    print([input_layer.name for input_layer in decoder_model.inputs])
    
  2. 避免手动拼接JSON:使用json.dumps()自动处理序列化,防止引号转义、格式不规范等问题。
  3. 响应解析适配:根据decoder_model的输出结构,调整result["predictions"]的索引,确保正确提取output_tokens和状态值。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 23:00:54