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

