SageMaker中使用JSONL数据批量预测时遇JSON解码错误
问题
我有一个已在SageMaker实时端点测试通过的PyTorch模型,现需切换至批量推理。使用JSONL数据,按照AWS文档配置批量转换任务,并自定义了inference.py,但在input_fn函数中执行json.loads(request_body)时,出现JSON解码错误:raise JSONDecodeError("Extra data", s, end)。该模型及JSON输入在实时端点测试正常,但批量任务报错。
相关代码
自定义inference.py
def model_fn(model_dir): .... def input_fn(request_body, request_content_type): data = json.loads(request_body) return data def predict_fn(data, model): ...
Lambda配置批量任务代码
response = client.create_transform_job( TransformJobName='some-job', ModelName='mypytorchmodel', ModelClientConfig={ 'InvocationsTimeoutInSeconds': 3600, 'InvocationsMaxRetries': 1 }, BatchStrategy='MultiRecord', TransformInput={ 'DataSource': { 'S3DataSource': { 'S3DataType': 'S3Prefix', 'S3Uri': 's3://inputpath' } }, 'ContentType': 'application/json', 'SplitType': 'Line' }, TransformOutput={ 'S3OutputPath': 's3://outputpath', 'Accept': 'application/json', 'AssembleWith': 'Line', }, TransformResources={ 'InstanceType': 'ml.g4dn.xlarge', 'InstanceCount': 1 } )
输入JSONL文件内容
{"input" : "input line one"} {"input" : "input line two"} {"input" : "input line three"} {"input" : "input line four"} {"input" : "input line five"} ...
解决方案
问题根源是BatchStrategy='MultiRecord'的配置:当设置为该值时,SageMaker会将多条JSONL记录拼接成一个请求体传给input_fn,而非单条记录。比如拼接后的内容类似{"input":"line1"}{"input":"line2"},这种格式不是合法的单一JSON对象/数组,自然无法被json.loads解析。
有两种解决方式:
- 修改批量任务配置:将
BatchStrategy改为SingleRecord。此时每条JSONL记录会单独传入input_fn,和实时端点调用逻辑一致,无需修改inference.py。 - 保留MultiRecord并修改inference.py:如果想利用批量打包提升效率,需要调整
input_fn来处理多条拼接的JSON记录:
import json def input_fn(request_body, request_content_type): # 按行分割请求体,逐个解析JSON记录 records = request_body.splitlines() data = [json.loads(line) for line in records if line.strip()] return data
同时要确保predict_fn能适配批量输入的数据结构(比如接收列表并返回对应结果列表)。
注意:TransformInput中的SplitType='Line'配置是正确的,它负责按行分割输入文件,但MultiRecord会将多行打包为一个请求。
内容的提问来源于stack exchange,提问作者arve
相关产品推荐
相关产品推荐

