如何在SageMaker Pipeline的TrainingStep中传递参数至Python训练脚本?
解决SageMaker Pipeline中训练脚本获取通道数据的问题
核心原因
SageMaker的TrainingInput通道不会自动以命令行参数形式传递给训练脚本,你之前用argparse获取--train和--validation参数的思路不成立——这些参数根本没被传入脚本启动命令,所以会返回None。
正确的两种实现方式
方式1:直接读取SageMaker默认通道挂载路径
SageMaker会自动把每个通道的数据集挂载到训练容器的固定路径:/opt/ml/input/data/<channel_name>,其中<channel_name>就是你在TrainingInput里定义的通道名(比如train、validation)。
修改xgboost_train.py脚本,直接使用该路径读取数据,无需argparse:
import pandas as pd import xgboost as xgb # 直接读取通道挂载路径下的数据 train_data_dir = "/opt/ml/input/data/train" val_data_dir = "/opt/ml/input/data/validation" # 示例:读取csv文件(根据你的预处理输出格式调整) train_df = pd.read_csv(f"{train_data_dir}/train.csv") val_df = pd.read_csv(f"{val_data_dir}/validation.csv") # 后续训练逻辑...
方式2:通过Estimator配置传递命令行参数
如果坚持要用argparse获取路径,需要在创建XGBoost Estimator时,手动把通道挂载路径作为命令行参数传入脚本:
- 修改Estimator配置,自定义启动命令传递参数:
from sagemaker.xgboost.estimator import XGBoost xgb_estimator = XGBoost( entry_point="xgboost_train.py", role=your_sagemaker_role, instance_count=1, instance_type="ml.m5.xlarge", framework_version="1.7-1", # 自定义命令,将通道路径作为参数传给脚本 command=["python3", "xgboost_train.py", "--train", "/opt/ml/input/data/train", "--validation", "/opt/ml/input/data/validation"] )
- 在
xgboost_train.py中用argparse接收参数:
import argparse import pandas as pd import xgboost as xgb parser = argparse.ArgumentParser() parser.add_argument("--train", type=str, required=True) parser.add_argument("--validation", type=str, required=True) args = parser.parse_args() # 使用传入的路径读取数据 train_df = pd.read_csv(f"{args.train}/train.csv") val_df = pd.read_csv(f"{args.validation}/validation.csv") # 后续训练逻辑...
额外注意事项
- 确保Pipeline的TrainingStep中定义的通道名称(如
train、validation)和脚本里使用的完全一致,包括大小写。 - 如果预处理输出是多个文件,需调整读取逻辑(比如用
os.listdir遍历目录下的文件)。
内容的提问来源于stack exchange,提问作者TYL
相关产品推荐
相关产品推荐

