You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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时,手动把通道挂载路径作为命令行参数传入脚本:

  1. 修改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"]
)
  1. 在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.20 06:53:14