SageMaker Pipeline中ProcessingStep.arguments引发不必要S3上传的解决方法
问题
我参考SageMaker Pipelines示例编写模型注册代码时遇到一个问题:代码中使用step_eval.arguments获取模型指标的S3路径时,会自动将包含evaluate.py的sourcedir.tar.gz上传到S3,但我只需要读取目标模型指标的S3路径,不需要这个上传行为。以下是我的模型注册代码:
import logging from sagemaker.drift_check_baselines import DriftCheckBaselines from sagemaker.model_metrics import ( MetricsSource, ModelMetrics, FileSource, ) from sagemaker.workflow.model_step import ModelStep logger = logging.getLogger() logger.setLevel(logging.INFO) logger.addHandler(logging.StreamHandler()) def define_register_step( step_name, model, step_eval, step_explainability, model_explainability_check_config, model_package_group_name, model_approval_status, ): ''' Register model step that will be conditionally executed ''' logger.info(f'\tStart of define_register_step') my_args = step_eval.arguments # TODO: 这会触发包含evaluation.py的sourcedir.tar.gz上传到S3 logger.info(f'\tstep_eval.arguments: {my_args}') s3_folder_path = my_args["ProcessingOutputConfig"]["Outputs"][0]["S3Output"]["S3Uri"] s3_uri="{}/evaluation.json".format(s3_folder_path) model_statistics=MetricsSource( s3_uri=s3_uri, content_type="application/json" ) model_metrics = ModelMetrics(model_statistics=model_statistics) drift_check_baselines = DriftCheckBaselines( explainability_constraints=MetricsSource( s3_uri=step_explainability.properties.BaselineUsedForDriftCheckConstraints, content_type="application/json", ), explainability_config_file=FileSource( s3_uri=model_explainability_check_config.monitoring_analysis_config_uri, content_type="application/json", ), ) step_args = model.register( content_types=["text/csv"], response_types=["text/csv"], inference_instances=["ml.t2.medium", "ml.m5.large"], transform_instances=["ml.m5.large"], model_package_group_name=model_package_group_name, approval_status=model_approval_status, model_metrics=model_metrics, drift_check_baselines=drift_check_baselines, ) step_register = ModelStep( name=step_name, step_args=step_args, ) return step_register
解决方案
问题根源在于直接访问step_eval.arguments:这个属性会返回ProcessingStep的完整参数配置,其中包含本地源代码路径,SageMaker Pipelines为了解析这些参数,会自动将依赖资源(如sourcedir.tar.gz)上传到S3。
正确做法是通过ProcessingStep的**properties属性**获取输出S3 URI,这是Pipeline内的动态引用,不会触发本地资源上传,仅在Pipeline运行时从步骤实际输出中读取路径。
修改代码核心部分
替换原来通过arguments获取路径的逻辑:
logger.info(f'\tStart of define_register_step') # 改用properties获取输出路径,避免不必要的资源上传 s3_folder_path = step_eval.properties.ProcessingOutputConfig.Outputs[0].S3Output.S3Uri logger.info(f'\tEvaluation output S3 path: {s3_folder_path}')
如果你的ProcessingStep在定义时给输出设置了名称(比如name="evaluation"),可以用名称索引输出,代码更清晰:
s3_folder_path = step_eval.properties.ProcessingOutputConfig.Outputs['evaluation'].S3Output.S3Uri
修改后的完整函数代码
import logging from sagemaker.drift_check_baselines import DriftCheckBaselines from sagemaker.model_metrics import ( MetricsSource, ModelMetrics, FileSource, ) from sagemaker.workflow.model_step import ModelStep logger = logging.getLogger() logger.setLevel(logging.INFO) logger.addHandler(logging.StreamHandler()) def define_register_step( step_name, model, step_eval, step_explainability, model_explainability_check_config, model_package_group_name, model_approval_status, ): ''' Register model step that will be conditionally executed ''' logger.info(f'\tStart of define_register_step') # 通过properties获取输出S3路径,避免触发不必要的资源上传 s3_folder_path = step_eval.properties.ProcessingOutputConfig.Outputs[0].S3Output.S3Uri logger.info(f'\tEvaluation output S3 path: {s3_folder_path}') s3_uri="{}/evaluation.json".format(s3_folder_path) model_statistics=MetricsSource( s3_uri=s3_uri, content_type="application/json" ) model_metrics = ModelMetrics(model_statistics=model_statistics) drift_check_baselines = DriftCheckBaselines( explainability_constraints=MetricsSource( s3_uri=step_explainability.properties.BaselineUsedForDriftCheckConstraints, content_type="application/json", ), explainability_config_file=FileSource( s3_uri=model_explainability_check_config.monitoring_analysis_config_uri, content_type="application/json", ), ) step_args = model.register( content_types=["text/csv"], response_types=["text/csv"], inference_instances=["ml.t2.medium", "ml.m5.large"], transform_instances=["ml.m5.large"], model_package_group_name=model_package_group_name, approval_status=model_approval_status, model_metrics=model_metrics, drift_check_baselines=drift_check_baselines, ) step_register = ModelStep( name=step_name, step_args=step_args, ) return step_register
原理说明
step_eval.arguments:属于步骤的静态参数定义,访问时会触发SageMaker Pipelines解析所有依赖资源(包括本地源代码),因此会自动上传sourcedir.tar.gz。step_eval.properties:属于步骤的动态输出属性引用,是Pipeline运行时的占位符,仅在Pipeline执行时获取步骤实际生成的输出路径,不会触发本地资源上传。
内容的提问来源于stack exchange,提问作者Francisco C
相关产品推荐
相关产品推荐

