AWS SageMaker端点调用400错误:请求格式排查求助
解决SageMaker TensorFlow端点请求格式问题
确定模型期望的请求格式
1. 查看模型输入签名
本地加载已保存的TensorFlow模型,检查其服务签名的输入参数名称:
import tensorflow as tf # 替换为你的模型本地路径 loaded_model = tf.saved_model.load("./your-saved-model-path") serving_signature = loaded_model.signatures["serving_default"] # 打印输入信息 for input_tensor in serving_signature.inputs: print(f"输入名称: {input_tensor.name}, 形状: {input_tensor.shape}, 类型: {input_tensor.dtype}")
输出会显示模型期望的输入参数名(比如inputs:0对应的参数名是inputs)、形状和数据类型,这是请求格式的核心依据。
2. 对应SageMaker端点的请求格式
SageMaker部署的TensorFlow模型默认使用TensorFlow Serving协议,请求格式分为两种常见类型:
- 如果模型输入是单值张量(比如单个字符串):
直接将数据放入instances数组,格式如下:
这和你在Jupyter实例中调用{ "instances": ["Testing model's prediction on this text"] }predictor.predict(['xxx'])时,SDK自动包装的格式一致。 - 如果模型输入是命名张量(比如输入参数名为
text):
需要在instances数组中传入键值对:{ "instances": [{"text": "Testing model's prediction on this text"}] }
测试正确的请求方式
在SageMaker Studio或其他Notebook中,用boto3调用端点的示例代码:
import boto3 import json sm_runtime = boto3.client("sagemaker-runtime") response = sm_runtime.invoke_endpoint( EndpointName="testing-2", ContentType="application/json", Body=json.dumps({"instances": ["Testing model's prediction on this text"]}) ) # 解析返回结果 predictions = json.loads(response["Body"].read().decode()) print(predictions)
关键说明
你本地用curl调用的是自定义服务(比如Flask),和SageMaker默认的TensorFlow Serving协议格式不同,因此不能直接沿用本地的请求结构。必须匹配模型签名定义的输入参数和SageMaker的服务协议格式。
内容的提问来源于stack exchange,提问作者Sarath
相关产品推荐
相关产品推荐

