使用Parquet文件的SageMaker XGBoost批量转换推理任务失败求助
解决SageMaker XGBoost批量推理Parquet格式不支持的问题
问题原因
SageMaker官方提供的XGBoost容器(如你使用的1.7-1版本),其默认推理脚本仅支持CSV、JSON等文本类格式输入,不直接支持Parquet二进制格式。训练阶段能通过Pipe模式处理Parquet是因为训练流程的输入处理逻辑和推理阶段完全独立,不能直接沿用。
解决方案
方案1:转换输入数据格式为CSV(快速解决)
将S3中的批量输入Parquet文件转换为CSV格式,然后修改批量推理任务的ContentType为text/csv:
修改后的推理任务代码:
create_batch = client.create_transform_job( TransformJobName=transformJobName, ModelName=Modelname, MaxConcurrentTransforms=0, MaxPayloadInMB=6, BatchStrategy='MultiRecord', TransformInput={ 'DataSource': { 'S3DataSource': { 'S3DataType': 'S3Prefix', 'S3Uri': batch_input_csv # 替换为CSV文件的S3路径 } }, "ContentType": "text/csv", # 修改为CSV格式 "CompressionType": "None", }, TransformOutput={ 'S3OutputPath': batch_output, 'AssembleWith': 'Line' }, TransformResources={ 'InstanceType': 'ml.m4.xlarge', 'InstanceCount': 1 } )
方案2:自定义推理脚本支持Parquet输入(长期适配)
通过自定义推理脚本扩展模型的输入格式支持,步骤如下:
- 编写自定义推理脚本(例如
inference.py),添加Parquet读取逻辑:
import pandas as pd import xgboost as xgb import os def model_fn(model_dir): # 加载训练好的XGBoost模型 model = xgb.Booster() model.load_model(os.path.join(model_dir, 'xgboost-model')) return model def input_fn(request_body, request_content_type): # 根据Content-Type处理输入数据 if request_content_type == 'application/x-parquet': # 读取Parquet数据 df = pd.read_parquet(request_body) return xgb.DMatrix(df) elif request_content_type == 'text/csv': # 保留默认CSV处理逻辑 df = pd.read_csv(request_body) return xgb.DMatrix(df) else: raise ValueError(f"Unsupported content type: {request_content_type}") def predict_fn(input_data, model): # 执行预测 predictions = model.predict(input_data) return predictions def output_fn(prediction, response_content_type): # 处理输出格式 return str(prediction)
- 在注册模型时,指定自定义脚本:
from sagemaker.xgboost.model import XGBoostModel # 创建带自定义脚本的模型 xgb_model = XGBoostModel( model_data='s3://your-bucket/path/to/model.tar.gz', # 训练输出的模型文件路径 role=role, entry_point='inference.py', # 指定自定义推理脚本 framework_version='1.7-1', sagemaker_session=sess ) # 注册模型到模型注册表 model_package = xgb_model.register( model_package_group_name='your-model-group', content_types=['application/x-parquet', 'text/csv'], response_types=['text/plain'] )
- 重新运行批量推理任务,此时
ContentType指定为application/x-parquet即可正常执行。
内容的提问来源于stack exchange,提问作者user3858193
相关产品推荐
相关产品推荐

