Amazon SageMaker训练的XGBoost模型如何加载并预测新数据
1 加载S3中已训练XGBoost模型到新Notebook并预测的代码示例
你训练完成的模型会自动打包为model.tar.gz存储在你训练代码中指定的output_path路径下,可直接调用SageMaker SDK加载使用,以下分批量预测(适合大体积新数据上传S3后批量生成结果)和实时端点预测(适合Notebook中快速调用少量数据预测)两种场景给出代码:
# 先导入基础依赖 import sagemaker import boto3 from sagemaker.predictor import CSVSerializer from sagemaker.xgboost.model import XGBoostModel role = sagemaker.get_execution_role() sess = sagemaker.Session() region = boto3.Session().region_name # 替换为你实际的模型文件路径,可在SageMaker控制台对应训练任务的详情页查到 model_data = "s3://innogy-bda-germany-dev-landing-dc3-retailpl/UPSELL/LIST/output/你的训练任务名称/model.tar.gz" # 初始化已训练的XGBoost模型对象 xgb_model = XGBoostModel( model_data=model_data, role=role, sagemaker_session=sess, framework_version="latest", entry_point=None )
1.1 批量预测场景(推荐,适合S3上传的批量新数据)
# 配置批量转换任务 transformer = xgb_model.transformer( instance_count=1, instance_type="ml.m5.xlarge", output_path="s3://innogy-bda-germany-dev-landing-dc3-retailpl/UPSELL/LIST/prediction_output/", # 预测结果存储路径 assemble_with="Line", accept="text/csv" ) # 新数据路径要求:csv文件不要包含目标列,列顺序和训练时的特征列完全一致,不要带表头 new_data_path = "s3://innogy-bda-germany-dev-landing-dc3-retailpl/UPSELL/LIST/new_data/" transformer.transform( new_data_path, content_type="text/csv", split_type="Line" ) # 等待任务完成后可直接到output_path下载结果 transformer.wait() print(f"预测结果已保存到:{transformer.output_path}")
1.2 实时端点预测场景(适合Notebook中临时调试)
# 部署实时预测端点 predictor = xgb_model.deploy( initial_instance_count=1, instance_type="ml.t2.medium", serializer=CSVSerializer() ) # 输入特征顺序和训练时的特征顺序完全一致即可,不需要传入目标列 test_features = [0.34, 2, 1, 0.76, ...] # 替换为你的实际特征值 prediction = predictor.predict(test_features) print(f"预测结果:{prediction}") # 测试完成后记得删除端点避免产生额外费用 # predictor.delete_endpoint()
2 SageMaker内置XGBoost训练无需提前剔除目标变量的原因
这是SageMaker内置XGBoost算法的固定输入规则决定的:训练/验证用的CSV文件要求第一列必须是目标变量,后续列是特征,算法底层会自动拆分第一列为标签y、剩余列为特征X,不需要用户手动处理目标列。
你给出的数据集拆分代码只是做了行维度的拆分,拆分后的训练/验证集仍然保留了「第一列为目标列、后续为特征」的结构,刚好符合内置算法的输入要求,所以不需要额外剔除目标列。
而预测阶段的输入没有目标列的要求,只要保证特征顺序和训练时完全一致即可,和训练阶段的逻辑互不冲突。
内容的提问来源于stack exchange,提问作者Patryk Kołakowski
相关产品推荐
相关产品推荐

