如何通过start_pipeline_execution向NotebookJobStep传动态参数?
无需重建SageMaker Pipeline即可向NotebookJobStep传递动态参数的方案
核心思路
利用SageMaker Pipeline的可配置参数机制,在Pipeline定义阶段预留参数占位符,触发执行时动态传入参数值,这些参数会自动传递给NotebookJobStep中的笔记本,全程无需重新构建或更新Pipeline。
具体实现步骤
1. 定义Pipeline参数(创建Pipeline时完成)
在构建Pipeline时,先声明需要动态传递的参数(比如输入路径、阈值、数据集名称等),使用ParameterString/ParameterInteger等类型:
from sagemaker.workflow.parameters import ParameterString, ParameterInteger # 定义动态参数,示例为输入数据路径和置信度阈值 input_data_param = ParameterString( name="InputDataPath", default_value="s3://default-bucket/default-dataset/" ) confidence_threshold_param = ParameterInteger( name="ConfidenceThreshold", default_value=0.8 )
2. 将参数关联到NotebookJobStep
创建NotebookJobStep时,通过notebook_params字段把Pipeline参数传递给笔记本,这些参数会以环境变量的形式注入到笔记本运行环境:
from sagemaker.workflow.steps import NotebookJobStep notebook_step = NotebookJobStep( name="RunDynamicNotebook", notebook_job_name="dynamic-notebook-job", notebook_s3_uri="s3://your-bucket/notebooks/your-notebook.ipynb", role="SageMakerExecutionRole", # 绑定Pipeline参数到笔记本环境变量 notebook_params={ "INPUT_DATA_PATH": input_data_param, "CONFIDENCE_THRESHOLD": confidence_threshold_param }, instance_type="ml.t3.medium" ) # 组装并创建Pipeline(只需执行一次) from sagemaker.workflow.pipeline import Pipeline pipeline = Pipeline( name="DynamicNotebookPipeline", parameters=[input_data_param, confidence_threshold_param], steps=[notebook_step] ) # 首次创建/更新Pipeline pipeline.upsert(role_arn="arn:aws:iam::123456789012:role/SageMakerExecutionRole")
3. 在笔记本中读取参数
在你的.ipynb笔记本里,通过读取环境变量获取传递的参数:
import os # 读取Pipeline传递的动态参数 input_data_path = os.environ.get("INPUT_DATA_PATH") confidence_threshold = float(os.environ.get("CONFIDENCE_THRESHOLD")) # 业务逻辑中使用参数 print(f"当前使用数据集路径: {input_data_path}") print(f"当前置信度阈值: {confidence_threshold}")
4. 触发Pipeline时动态传入参数
使用boto3调用start_pipeline_execution时,通过parameters参数传入本次执行的具体值,无需修改Pipeline本身:
import boto3 sagemaker_client = boto3.client("sagemaker") # 每次触发传入不同参数值 response = sagemaker_client.start_pipeline_execution( PipelineName="DynamicNotebookPipeline", PipelineExecutionDisplayName="Execution-20240520-Batch", Parameters={ "InputDataPath": "s3://your-bucket/data/batch-20240520/", "ConfidenceThreshold": "0.9" # 注意:所有参数值需以字符串传递 } ) print(f"Pipeline执行ARN: {response['PipelineExecutionArn']}")
关键注意事项
- 触发时传递的所有参数值必须是字符串格式,数值类型需转成字符串后传入,笔记本中再做类型转换。
- 若需传递复杂结构(如JSON),可先序列化为字符串,在笔记本中反序列化使用。
- 首次创建Pipeline后,后续触发只需修改
Parameters中的值,无需重新调用pipeline.upsert()。
内容的提问来源于stack exchange,提问作者Eduardus Bagaskara
相关产品推荐
相关产品推荐

