如何在SageMaker PyTorch TorchServe端点启用服务端批处理?
在Amazon SageMaker PyTorch TorchServe端点启用服务端批处理的方法
1. 配置TorchServe批处理参数
在模型归档包中添加config.properties文件,设置核心批处理参数:
# 开启批处理 enable_batch_inference=true # 单批次最大请求数(根据模型和实例资源调整) batch_size=8 # 等待凑齐批次的最长延迟(毫秒,超时则直接处理现有请求) max_batch_delay=100 # 可选:批次处理超时时间 batch_timeout=5000
2. 打包模型时包含配置文件
使用torch-model-archiver打包模型时,通过--extra-files参数引入上述配置文件:
torch-model-archiver --model-name my_model --version 1.0 --model-file model.py --serialized-file model.pth --handler handler.py --extra-files config.properties
将打包好的.mar文件上传至S3存储桶。
3. SageMaker端点部署配置
通过SageMaker Python SDK部署时,显式开启批处理支持并设置相关环境变量(可覆盖config.properties中的参数):
from sagemaker.pytorch.model import PyTorchModel pytorch_model = PyTorchModel( model_data="s3://your-bucket/model.tar.gz", role="your-sagemaker-iam-role", framework_version="2.1", py_version="py310", env={ "SAGEMAKER_TORCHSERVE_ENABLE_BATCHING": "true", "SAGEMAKER_TORCHSERVE_BATCH_SIZE": "8", "SAGEMAKER_TORCHSERVE_MAX_BATCH_DELAY": "100" } ) predictor = pytorch_model.deploy( initial_instance_count=1, instance_type="ml.g4dn.xlarge" # 建议用GPU实例提升批处理效率 )
4. 适配Handler支持批处理
确保自定义Handler的handle方法能处理批量输入数据,示例如下:
def handle(data, context): # 解析批量输入 inputs = [item["input"] for item in data] # 模型批量推理(需模型本身支持批量输入) outputs = model(inputs) # 整理批量输出格式 return [{"output": out} for out in outputs]
5. 测试批处理请求
发送包含多个样本的批量请求,格式需匹配Handler预期:
# 构造8条测试数据 test_batch = [{"input": f"sample_{i}"} for i in range(8)] # 发送批量推理请求 response = predictor.predict(test_batch)
注意事项
- 实例资源匹配:
batch_size需根据实例的内存/GPU显存调整,避免OOM; - 框架版本:建议使用PyTorch 1.13及以上版本,确保TorchServe批处理功能兼容;
- 延迟权衡:
max_batch_delay设置过小可能无法凑齐批次,过大则增加请求延迟,需根据业务场景调整。
内容的提问来源于stack exchange,提问作者Francesco Pochetti
相关产品推荐
相关产品推荐

