diff --git a/sagemaker-train/src/sagemaker/train/dpo_trainer.py b/sagemaker-train/src/sagemaker/train/dpo_trainer.py index ef6997d7e2..8f7c7287a2 100644 --- a/sagemaker-train/src/sagemaker/train/dpo_trainer.py +++ b/sagemaker-train/src/sagemaker/train/dpo_trainer.py @@ -5,6 +5,7 @@ from sagemaker.train.common import TrainingType, CustomizationTechnique, JOB_TYPE from sagemaker.core.resources import TrainingJob, ModelPackageGroup, ModelPackage from sagemaker.core.shapes import VpcConfig +from sagemaker.core.workflow.pipeline_context import PipelineSession from sagemaker.train.defaults import TrainDefaults from sagemaker.train.utils import _get_unique_name, _get_jumpstart_tags from sagemaker.train.configs import StoppingCondition @@ -369,6 +370,14 @@ def train(self, if self.stopping_condition is not None: create_args["stopping_condition"] = self.stopping_condition + # If running within a PipelineSession, intercept the request and store + # step arguments instead of launching a training job. + # This must come before data path validation since in pipeline mode + # the data path may be a pipeline parameter that doesn't exist yet. + if isinstance(sagemaker_session, PipelineSession): + sagemaker_session._intercept_create_request(create_args, None, "train") + return sagemaker_session.context + # Validate data paths exist before submission effective_training = training_dataset or self.training_dataset effective_validation = validation_dataset or self.validation_dataset diff --git a/sagemaker-train/src/sagemaker/train/rlaif_trainer.py b/sagemaker-train/src/sagemaker/train/rlaif_trainer.py index a3060cec0f..ecee211ba0 100644 --- a/sagemaker-train/src/sagemaker/train/rlaif_trainer.py +++ b/sagemaker-train/src/sagemaker/train/rlaif_trainer.py @@ -4,6 +4,7 @@ from sagemaker.train.common import TrainingType, CustomizationTechnique, JOB_TYPE from sagemaker.core.resources import TrainingJob, ModelPackageGroup, MlflowTrackingServer, ModelPackage from sagemaker.core.shapes import VpcConfig +from sagemaker.core.workflow.pipeline_context import PipelineSession from sagemaker.train.defaults import TrainDefaults from sagemaker.train.utils import _get_unique_name, _get_jumpstart_tags from sagemaker.train.common_utils.recipe_utils import _get_hub_content_metadata @@ -335,6 +336,14 @@ def train(self, training_dataset: Optional[Union[str, DataSet]] = None, validati if self.stopping_condition is not None: create_args["stopping_condition"] = self.stopping_condition + # If running within a PipelineSession, intercept the request and store + # step arguments instead of launching a training job. + # This must come before data path validation since in pipeline mode + # the data path may be a pipeline parameter that doesn't exist yet. + if isinstance(sagemaker_session, PipelineSession): + sagemaker_session._intercept_create_request(create_args, None, "train") + return sagemaker_session.context + # Validate data paths exist before submission effective_training = training_dataset or self.training_dataset effective_validation = validation_dataset or self.validation_dataset diff --git a/sagemaker-train/src/sagemaker/train/rlvr_trainer.py b/sagemaker-train/src/sagemaker/train/rlvr_trainer.py index 5f03cb5b8c..470854c4ac 100644 --- a/sagemaker-train/src/sagemaker/train/rlvr_trainer.py +++ b/sagemaker-train/src/sagemaker/train/rlvr_trainer.py @@ -6,6 +6,7 @@ from sagemaker.train.common import TrainingType, CustomizationTechnique, JOB_TYPE from sagemaker.core.resources import TrainingJob, ModelPackageGroup, MlflowTrackingServer, ModelPackage from sagemaker.core.shapes import VpcConfig +from sagemaker.core.workflow.pipeline_context import PipelineSession from sagemaker.train.defaults import TrainDefaults from sagemaker.train.utils import _get_unique_name, _get_jumpstart_tags from sagemaker.ai_registry.dataset import DataSet @@ -555,6 +556,14 @@ def train(self, training_dataset: Optional[Union[str, DataSet]] = None, if self.stopping_condition is not None: create_args["stopping_condition"] = self.stopping_condition + # If running within a PipelineSession, intercept the request and store + # step arguments instead of launching a training job. + # This must come before data path validation since in pipeline mode + # the data path may be a pipeline parameter that doesn't exist yet. + if isinstance(sagemaker_session, PipelineSession): + sagemaker_session._intercept_create_request(create_args, None, "train") + return sagemaker_session.context + # Validate data paths exist before submission effective_training = training_dataset or self.training_dataset effective_validation = validation_dataset or self.validation_dataset diff --git a/sagemaker-train/src/sagemaker/train/sft_trainer.py b/sagemaker-train/src/sagemaker/train/sft_trainer.py index eb06d23905..a3ac08e827 100644 --- a/sagemaker-train/src/sagemaker/train/sft_trainer.py +++ b/sagemaker-train/src/sagemaker/train/sft_trainer.py @@ -4,6 +4,7 @@ from sagemaker.train.common import TrainingType, CustomizationTechnique, JOB_TYPE from sagemaker.core.resources import TrainingJob, ModelPackageGroup, ModelPackage from sagemaker.core.shapes import VpcConfig +from sagemaker.core.workflow.pipeline_context import PipelineSession from sagemaker.train.defaults import TrainDefaults from sagemaker.train.utils import _get_unique_name, _get_jumpstart_tags from sagemaker.ai_registry.dataset import DataSet @@ -437,6 +438,14 @@ def train(self, training_dataset: Optional[Union[str, DataSet]] = None, validati if self.stopping_condition is not None: create_args["stopping_condition"] = self.stopping_condition + # If running within a PipelineSession, intercept the request and store + # step arguments instead of launching a training job. + # This must come before data path validation since in pipeline mode + # the data path may be a pipeline parameter that doesn't exist yet. + if isinstance(sagemaker_session, PipelineSession): + sagemaker_session._intercept_create_request(create_args, None, "train") + return sagemaker_session.context + # Validate data paths exist before submission effective_training = training_dataset or self.training_dataset effective_validation = validation_dataset or self.validation_dataset diff --git a/sagemaker-train/tests/unit/train/test_dpo_trainer.py b/sagemaker-train/tests/unit/train/test_dpo_trainer.py index 8ef48041a3..b679e54618 100644 --- a/sagemaker-train/tests/unit/train/test_dpo_trainer.py +++ b/sagemaker-train/tests/unit/train/test_dpo_trainer.py @@ -750,3 +750,72 @@ def test_list_supported_models(self, mock_list): mock_list.assert_called_once_with( recipe_type="FineTuning", technique="DPO", session=None ) + +class TestDPOTrainerPipelineSession: + """Test DPOTrainer behavior when PipelineSession is used. + + Ref: https://github.com/aws/sagemaker-python-sdk/issues/6163 + """ + + @patch('sagemaker.train.dpo_trainer._create_model_package_config') + @patch('sagemaker.train.dpo_trainer._create_mlflow_config') + @patch('sagemaker.train.dpo_trainer._create_output_config') + @patch('sagemaker.train.dpo_trainer._create_serverless_config') + @patch('sagemaker.train.dpo_trainer._convert_input_data_to_channels') + @patch('sagemaker.train.dpo_trainer._create_input_data_config') + @patch('sagemaker.train.dpo_trainer._get_unique_name') + @patch('sagemaker.train.dpo_trainer.TrainDefaults.get_role') + @patch('sagemaker.train.dpo_trainer.TrainDefaults.get_sagemaker_session') + @patch('sagemaker.train.dpo_trainer._validate_and_resolve_model_package_group') + @patch('sagemaker.train.dpo_trainer._get_fine_tuning_options_and_model_arn') + @patch('sagemaker.train.dpo_trainer._resolve_model_and_name') + @patch('sagemaker.train.common_utils.finetune_utils._get_beta_session') + @patch('sagemaker.core.resources.TrainingJob.create') + def test_train_with_pipeline_session_does_not_launch_job( + self, mock_training_job_create, mock_beta_session, mock_resolve_model, + mock_finetuning_options, mock_validate_group, mock_get_session, mock_get_role, + mock_unique_name, mock_input_config, mock_convert_channels, + mock_serverless_config, mock_output_config, mock_mlflow_config, mock_model_package_config, + ): + """When PipelineSession is passed, _intercept_create_request traps the args.""" + from sagemaker.train.dpo_trainer import DPOTrainer + from sagemaker.core.workflow.pipeline_context import PipelineSession, _JobStepArguments + + pipeline_session = Mock(spec=PipelineSession) + pipeline_session.boto_session = Mock() + pipeline_session.boto_session.region_name = "us-west-2" + + step_args = _JobStepArguments("train", {"training_job_name": "test-dpo-job-001"}) + pipeline_session._intercept_create_request.return_value = None + pipeline_session.context = step_args + mock_get_session.return_value = pipeline_session + + mock_resolve_model.return_value = ("test-model", "resolved-model-name") + mock_hyperparams = Mock() + mock_hyperparams.to_dict.return_value = {"param1": "value1"} + mock_hyperparams._specs = {"param1": {"type": "string"}} + mock_hyperparams._user_set = set() + mock_finetuning_options.return_value = (mock_hyperparams, "arn:aws:sagemaker:us-west-2:123456789012:model/test", False) + mock_validate_group.return_value = "test-group" + mock_get_role.return_value = "arn:aws:iam::123456789012:role/Role" + mock_unique_name.return_value = "test-dpo-job-001" + mock_input_config.return_value = {"train": "s3://bucket/data"} + mock_convert_channels.return_value = [{"ChannelName": "train"}] + mock_serverless_config.return_value = {"BaseModelArn": "arn:model"} + mock_output_config.return_value = {"S3OutputPath": "s3://bucket/output"} + mock_mlflow_config.return_value = None + mock_model_package_config.return_value = None + mock_beta_session.return_value = pipeline_session + + trainer = DPOTrainer(model="test-model", training_dataset="s3://bucket/data", model_package_group="test-group", sagemaker_session=pipeline_session) + trainer._model_arn = "arn:aws:sagemaker:us-west-2:123456789012:model/test" + trainer._model_name = "test-model" + trainer.accept_eula = True + trainer.hyperparameters = mock_hyperparams + + result = trainer.train() + + mock_training_job_create.assert_not_called() + pipeline_session._intercept_create_request.assert_called_once() + assert pipeline_session._intercept_create_request.call_args[0][2] == "train" + assert result == step_args diff --git a/sagemaker-train/tests/unit/train/test_rlaif_trainer.py b/sagemaker-train/tests/unit/train/test_rlaif_trainer.py index 9851f87c6f..fc8ccbd9cd 100644 --- a/sagemaker-train/tests/unit/train/test_rlaif_trainer.py +++ b/sagemaker-train/tests/unit/train/test_rlaif_trainer.py @@ -872,3 +872,72 @@ def test_list_supported_models(self, mock_list): mock_list.assert_called_once_with( recipe_type="FineTuning", technique="RLAIF", session=None ) + +class TestRLAIFTrainerPipelineSession: + """Test RLAIFTrainer behavior when PipelineSession is used. + + Ref: https://github.com/aws/sagemaker-python-sdk/issues/6163 + """ + + @patch('sagemaker.train.rlaif_trainer._create_model_package_config') + @patch('sagemaker.train.rlaif_trainer._create_mlflow_config') + @patch('sagemaker.train.rlaif_trainer._create_output_config') + @patch('sagemaker.train.rlaif_trainer._create_serverless_config') + @patch('sagemaker.train.rlaif_trainer._convert_input_data_to_channels') + @patch('sagemaker.train.rlaif_trainer._create_input_data_config') + @patch('sagemaker.train.rlaif_trainer._get_unique_name') + @patch('sagemaker.train.rlaif_trainer.TrainDefaults.get_role') + @patch('sagemaker.train.rlaif_trainer.TrainDefaults.get_sagemaker_session') + @patch('sagemaker.train.rlaif_trainer._validate_and_resolve_model_package_group') + @patch('sagemaker.train.rlaif_trainer._get_fine_tuning_options_and_model_arn') + @patch('sagemaker.train.rlaif_trainer._resolve_model_and_name') + @patch('sagemaker.train.common_utils.finetune_utils._get_beta_session') + @patch('sagemaker.core.resources.TrainingJob.create') + def test_train_with_pipeline_session_does_not_launch_job( + self, mock_training_job_create, mock_beta_session, mock_resolve_model, + mock_finetuning_options, mock_validate_group, mock_get_session, mock_get_role, + mock_unique_name, mock_input_config, mock_convert_channels, + mock_serverless_config, mock_output_config, mock_mlflow_config, mock_model_package_config, + ): + """When PipelineSession is passed, _intercept_create_request traps the args.""" + from sagemaker.train.rlaif_trainer import RLAIFTrainer + from sagemaker.core.workflow.pipeline_context import PipelineSession, _JobStepArguments + + pipeline_session = Mock(spec=PipelineSession) + pipeline_session.boto_session = Mock() + pipeline_session.boto_session.region_name = "us-west-2" + + step_args = _JobStepArguments("train", {"training_job_name": "test-rlaif-job-001"}) + pipeline_session._intercept_create_request.return_value = None + pipeline_session.context = step_args + mock_get_session.return_value = pipeline_session + + mock_resolve_model.return_value = ("test-model", "resolved-model-name") + mock_hyperparams = Mock() + mock_hyperparams.to_dict.return_value = {"param1": "value1"} + mock_hyperparams._specs = {"param1": {"type": "string"}} + mock_hyperparams._user_set = set() + mock_finetuning_options.return_value = (mock_hyperparams, "arn:aws:sagemaker:us-west-2:123456789012:model/test", False) + mock_validate_group.return_value = "test-group" + mock_get_role.return_value = "arn:aws:iam::123456789012:role/Role" + mock_unique_name.return_value = "test-rlaif-job-001" + mock_input_config.return_value = {"train": "s3://bucket/data"} + mock_convert_channels.return_value = [{"ChannelName": "train"}] + mock_serverless_config.return_value = {"BaseModelArn": "arn:model"} + mock_output_config.return_value = {"S3OutputPath": "s3://bucket/output"} + mock_mlflow_config.return_value = None + mock_model_package_config.return_value = None + mock_beta_session.return_value = pipeline_session + + trainer = RLAIFTrainer(model="test-model", training_dataset="s3://bucket/data", model_package_group="test-group", sagemaker_session=pipeline_session) + trainer._model_arn = "arn:aws:sagemaker:us-west-2:123456789012:model/test" + trainer._model_name = "test-model" + trainer.accept_eula = True + trainer.hyperparameters = mock_hyperparams + + result = trainer.train() + + mock_training_job_create.assert_not_called() + pipeline_session._intercept_create_request.assert_called_once() + assert pipeline_session._intercept_create_request.call_args[0][2] == "train" + assert result == step_args diff --git a/sagemaker-train/tests/unit/train/test_rlvr_trainer.py b/sagemaker-train/tests/unit/train/test_rlvr_trainer.py index 25839df202..dc7b56d29f 100644 --- a/sagemaker-train/tests/unit/train/test_rlvr_trainer.py +++ b/sagemaker-train/tests/unit/train/test_rlvr_trainer.py @@ -775,3 +775,71 @@ def test_list_supported_models(self, mock_list): mock_list.assert_called_once_with( recipe_type="FineTuning", technique="RLVR", session=None ) + +class TestRLVRTrainerPipelineSession: + """Test RLVRTrainer behavior when PipelineSession is used. + + Ref: https://github.com/aws/sagemaker-python-sdk/issues/6163 + """ + + @patch('sagemaker.train.rlvr_trainer._create_model_package_config') + @patch('sagemaker.train.rlvr_trainer._create_mlflow_config') + @patch('sagemaker.train.rlvr_trainer._create_output_config') + @patch('sagemaker.train.rlvr_trainer._convert_input_data_to_channels') + @patch('sagemaker.train.rlvr_trainer._create_input_data_config') + @patch('sagemaker.train.rlvr_trainer._get_unique_name') + @patch('sagemaker.train.rlvr_trainer.TrainDefaults.get_role') + @patch('sagemaker.train.rlvr_trainer.TrainDefaults.get_sagemaker_session') + @patch('sagemaker.train.rlvr_trainer._validate_and_resolve_model_package_group') + @patch('sagemaker.train.rlvr_trainer._get_fine_tuning_options_and_model_arn') + @patch('sagemaker.train.rlvr_trainer._resolve_model_and_name') + @patch('sagemaker.train.common_utils.finetune_utils._get_beta_session') + @patch('sagemaker.core.resources.TrainingJob.create') + def test_train_with_pipeline_session_does_not_launch_job( + self, mock_training_job_create, mock_beta_session, mock_resolve_model, + mock_finetuning_options, mock_validate_group, mock_get_session, mock_get_role, + mock_unique_name, mock_input_config, mock_convert_channels, + mock_output_config, mock_mlflow_config, mock_model_package_config, + ): + """When PipelineSession is passed, _intercept_create_request traps the args.""" + from sagemaker.train.rlvr_trainer import RLVRTrainer + from sagemaker.core.workflow.pipeline_context import PipelineSession, _JobStepArguments + + pipeline_session = Mock(spec=PipelineSession) + pipeline_session.boto_session = Mock() + pipeline_session.boto_session.region_name = "us-west-2" + + step_args = _JobStepArguments("train", {"training_job_name": "test-rlvr-job-001"}) + pipeline_session._intercept_create_request.return_value = None + pipeline_session.context = step_args + mock_get_session.return_value = pipeline_session + + mock_resolve_model.return_value = ("test-model", "resolved-model-name") + mock_hyperparams = Mock() + mock_hyperparams.to_dict.return_value = {"param1": "value1"} + mock_hyperparams._specs = {"param1": {"type": "string"}} + mock_hyperparams._user_set = set() + mock_finetuning_options.return_value = (mock_hyperparams, "arn:aws:sagemaker:us-west-2:123456789012:model/test", False) + mock_validate_group.return_value = "test-group" + mock_get_role.return_value = "arn:aws:iam::123456789012:role/Role" + mock_unique_name.return_value = "test-rlvr-job-001" + mock_input_config.return_value = {"train": "s3://bucket/data"} + mock_convert_channels.return_value = [{"ChannelName": "train"}] + mock_output_config.return_value = {"S3OutputPath": "s3://bucket/output"} + mock_mlflow_config.return_value = None + mock_model_package_config.return_value = None + mock_beta_session.return_value = pipeline_session + + trainer = RLVRTrainer(model="test-model", training_dataset="s3://bucket/data", model_package_group="test-group", sagemaker_session=pipeline_session) + trainer._model_arn = "arn:aws:sagemaker:us-west-2:123456789012:model/test" + trainer._model_name = "test-model" + trainer.accept_eula = True + trainer.hyperparameters = mock_hyperparams + trainer.custom_reward_function = None + + result = trainer.train() + + mock_training_job_create.assert_not_called() + pipeline_session._intercept_create_request.assert_called_once() + assert pipeline_session._intercept_create_request.call_args[0][2] == "train" + assert result == step_args diff --git a/sagemaker-train/tests/unit/train/test_sft_trainer.py b/sagemaker-train/tests/unit/train/test_sft_trainer.py index 1239bbd6d8..82cacfac45 100644 --- a/sagemaker-train/tests/unit/train/test_sft_trainer.py +++ b/sagemaker-train/tests/unit/train/test_sft_trainer.py @@ -1578,3 +1578,75 @@ def test_list_supported_models_passes_session(self, mock_list): mock_list.assert_called_once_with( recipe_type="FineTuning", technique="SFT", session=session ) + +class TestSFTTrainerPipelineSession: + """Test SFTTrainer behavior when PipelineSession is used. + + Ref: https://github.com/aws/sagemaker-python-sdk/issues/6163 + """ + + @patch('sagemaker.train.sft_trainer._validate_hyperparameter_values') + @patch('sagemaker.train.sft_trainer._create_model_package_config') + @patch('sagemaker.train.sft_trainer._create_mlflow_config') + @patch('sagemaker.train.sft_trainer._create_output_config') + @patch('sagemaker.train.sft_trainer._create_serverless_config') + @patch('sagemaker.train.sft_trainer._convert_input_data_to_channels') + @patch('sagemaker.train.sft_trainer._create_input_data_config') + @patch('sagemaker.train.sft_trainer._get_jumpstart_tags') + @patch('sagemaker.train.sft_trainer._get_unique_name') + @patch('sagemaker.train.sft_trainer.TrainDefaults.get_role') + @patch('sagemaker.train.sft_trainer.TrainDefaults.get_sagemaker_session') + @patch('sagemaker.train.sft_trainer._validate_and_resolve_model_package_group') + @patch('sagemaker.train.sft_trainer._get_fine_tuning_options_and_model_arn') + @patch('sagemaker.train.sft_trainer._resolve_model_and_name') + @patch('sagemaker.train.common_utils.finetune_utils._get_beta_session') + @patch('sagemaker.core.resources.TrainingJob.create') + def test_train_with_pipeline_session_does_not_launch_job( + self, mock_training_job_create, mock_beta_session, mock_resolve_model, + mock_finetuning_options, mock_validate_group, mock_get_session, mock_get_role, + mock_unique_name, mock_get_tags, mock_input_config, mock_convert_channels, + mock_serverless_config, mock_output_config, mock_mlflow_config, mock_model_package_config, + mock_validate_hp, + ): + """When PipelineSession is passed, _intercept_create_request traps the args.""" + from sagemaker.core.workflow.pipeline_context import PipelineSession, _JobStepArguments + + pipeline_session = Mock(spec=PipelineSession) + pipeline_session.boto_session = Mock() + pipeline_session.boto_session.region_name = "us-west-2" + + step_args = _JobStepArguments("train", {"training_job_name": "test-sft-job-001"}) + pipeline_session._intercept_create_request.return_value = None + pipeline_session.context = step_args + mock_get_session.return_value = pipeline_session + + mock_resolve_model.return_value = ("test-model", "resolved-model-name") + mock_hyperparams = Mock() + mock_hyperparams.to_dict.return_value = {"param1": "value1"} + mock_hyperparams._specs = {"param1": {"type": "string"}} + mock_hyperparams._user_set = set() + mock_finetuning_options.return_value = (mock_hyperparams, "arn:aws:sagemaker:us-west-2:123456789012:model/test", False) + mock_validate_group.return_value = "test-group" + mock_get_role.return_value = "arn:aws:iam::123456789012:role/Role" + mock_unique_name.return_value = "test-sft-job-001" + mock_get_tags.return_value = [] + mock_input_config.return_value = {"train": "s3://bucket/data"} + mock_convert_channels.return_value = [{"ChannelName": "train"}] + mock_serverless_config.return_value = {"BaseModelArn": "arn:model"} + mock_output_config.return_value = {"S3OutputPath": "s3://bucket/output"} + mock_mlflow_config.return_value = None + mock_model_package_config.return_value = None + mock_beta_session.return_value = pipeline_session + + trainer = SFTTrainer(model="test-model", training_dataset="s3://bucket/data", model_package_group="test-group", sagemaker_session=pipeline_session) + trainer._model_arn = "arn:aws:sagemaker:us-west-2:123456789012:model/test" + trainer._model_name = "test-model" + trainer.accept_eula = True + trainer.hyperparameters = mock_hyperparams + + result = trainer.train() + + mock_training_job_create.assert_not_called() + pipeline_session._intercept_create_request.assert_called_once() + assert pipeline_session._intercept_create_request.call_args[0][2] == "train" + assert result == step_args