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

如何以CSV文件为请求体调用SageMaker上部署的TensorFlow模型

我来一步步帮你搞定这个问题——不管是实现符合要求的serving_input_fn,还是后续用CSV请求调用SageMaker上的TensorFlow模型,都给你讲得明明白白。

第一步:实现适配CSV输入的serving_input_fn

当你用CSV作为请求体调用SageMaker的TensorFlow模型时,服务会把整个CSV请求体以字符串的形式传入模型,所以我们的serving_input_fn要做的就是接收这个字符串,把它解析成模型能处理的张量格式。

下面是完整的实现代码,我会逐行给你解释:

import tensorflow as tf

def serving_input_fn(hyperparameters):
    # 1. 定义接收CSV请求的占位符
    # shape=[None]表示可以接收任意数量的样本(每行CSV对应一个样本)
    csv_placeholder = tf.compat.v1.placeholder(tf.string, shape=[None])
    
    # 2. 预处理CSV数据:把字符串转换成模型需要的数值张量
    # 第一步:按逗号分割每行的CSV数据
    split_rows = tf.string_split(csv_placeholder, sep=',')
    # 第二步:把稀疏张量转换成稠密张量,用'0'填充缺失值(根据你的需求调整)
    dense_rows = tf.sparse.to_dense(split_rows, default_value='0')
    # 第三步:把字符串转换成模型需要的数值类型(比如float32,和训练时一致)
    feature_tensor = tf.string_to_number(dense_rows, out_type=tf.float32)
    
    # 3. 整理成模型期望的特征字典
    # 这里的键'features'要和你训练模型时定义的输入特征名完全匹配
    features = {'features': feature_tensor}
    
    # 返回ServingInputReceiver,关联特征张量和输入占位符
    return tf.estimator.export.ServingInputReceiver(features, {'csv_input': csv_placeholder})

关键细节提醒:

  • 如果你的CSV用的不是逗号分隔,把sep=','改成你实际的分隔符(比如';');
  • 如果训练时对特征做了归一化、编码等预处理,一定要在这个函数里复现相同的逻辑,否则模型预测结果会出错;
  • 特征字典的键要和模型的输入层名称完全一致,比如你模型输入是tf.keras.layers.Input(shape=(5,), name='user_features'),那这里就要写成{'user_features': feature_tensor}。
第二步:导出模型并部署到SageMaker

有了serving_input_fn之后,你需要用它来导出符合SageMaker要求的SavedModel:

# 假设你已经有训练好的estimator
estimator.export_savedmodel(
    export_dir_base='./exported_model',
    serving_input_receiver_fn=lambda: serving_input_fn(hyperparameters)
)

把导出的模型文件夹打包成.tar.gz,上传到S3的某个路径,然后在SageMaker控制台或者用SDK创建模型端点——这部分和普通模型部署流程一致,就不多啰嗦了。

第三步:用CSV请求体调用模型

部署好端点之后,你可以用两种常见方式调用:

方式1:用Python的requests库直接调用

import requests

# 替换成你的SageMaker端点URL
endpoint_url = "https://runtime.sagemaker.us-east-1.amazonaws.com/endpoints/your-endpoint-name/invocations"
# 你的CSV数据,每行一个样本(不要带表头,除非你在serving_input_fn里处理了表头)
csv_body = """1.2,3.4,5.6,7.8,9.0
2.3,4.5,6.7,8.9,0.1
3.4,5.6,7.8,9.0,1.2"""

# 设置请求头为text/csv,告诉SageMaker我们发送的是CSV数据
headers = {'Content-Type': 'text/csv'}
response = requests.post(endpoint_url, data=csv_body, headers=headers)

# 解析预测结果
predictions = response.json()
print(predictions)

方式2:用AWS SDK(boto3)调用

import boto3
import json

sagemaker_runtime = boto3.client('sagemaker-runtime')
endpoint_name = "your-endpoint-name"
csv_body = """1.2,3.4,5.6,7.8,9.0
2.3,4.5,6.7,8.9,0.1"""

response = sagemaker_runtime.invoke_endpoint(
    EndpointName=endpoint_name,
    ContentType='text/csv',
    Body=csv_body
)

# 读取并解析响应
result = json.loads(response['Body'].read().decode('utf-8'))
print(result)

额外注意点:

  • 如果你的CSV包含字符串类型的特征(比如分类变量),你需要在serving_input_fn里添加对应的编码逻辑(比如用VocabularyLookup),和训练时的处理保持一致;
  • 测试的时候可以先传单个样本的CSV,确认没问题后再批量传多个样本。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 09:43:22