Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 9 additions & 0 deletions sagemaker-train/src/sagemaker/train/dpo_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
9 changes: 9 additions & 0 deletions sagemaker-train/src/sagemaker/train/rlaif_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
9 changes: 9 additions & 0 deletions sagemaker-train/src/sagemaker/train/rlvr_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
9 changes: 9 additions & 0 deletions sagemaker-train/src/sagemaker/train/sft_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
69 changes: 69 additions & 0 deletions sagemaker-train/tests/unit/train/test_dpo_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
69 changes: 69 additions & 0 deletions sagemaker-train/tests/unit/train/test_rlaif_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
68 changes: 68 additions & 0 deletions sagemaker-train/tests/unit/train/test_rlvr_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Loading
Loading