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
16 changes: 12 additions & 4 deletions sagemaker-train/src/sagemaker/train/model_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
from sagemaker.core.shapes import AlgorithmSpecification, ModelPackageConfig
from sagemaker.core.utils.utils import serialize
from sagemaker.core.apiutils._boto_functions import to_pascal_case
from sagemaker.core.common_utils import name_from_base

from pydantic import BaseModel, ConfigDict, PrivateAttr, validate_call
from sagemaker.core.config.config_schema import (
Expand Down Expand Up @@ -75,7 +76,6 @@
from sagemaker.train.distributed import Torchrun, DistributedConfig
from sagemaker.train.utils import (
_default_s3_uri,
_get_unique_name,
_is_valid_path,
_is_valid_s3_uri,
safe_serialize,
Expand Down Expand Up @@ -619,7 +619,13 @@ def _create_training_job_args(
Dict[str, Any]: The training job arguments.
"""
self._populate_intelligent_defaults()
current_training_job_name = _get_unique_name(self.base_job_name)
# Use the shared name_from_base timestamp format so that
# sagemaker.core.common_utils.base_from_name can strip it back to base_job_name.
# trim_request_dict relies on base_from_name to preserve the prefix when a pipeline
# sets PipelineDefinitionConfig(use_custom_job_prefix=True) (issues #5776, #6299).
# Underscores are replaced because SageMaker job names disallow them; _get_unique_name
# used to do this and name_from_base does not.
current_training_job_name = name_from_base(self.base_job_name.replace("_", "-"))
input_data_key_prefix = f"{self.base_job_name}/{current_training_job_name}/input"

final_input_data_config = self.input_data_config.copy() if self.input_data_config else []
Expand Down Expand Up @@ -814,8 +820,10 @@ def _create_training_job_args(
training_request["model_package_config"] = self.model_package_config

if boto3 or isinstance(self.sagemaker_session, PipelineSession):
if isinstance(self.sagemaker_session, PipelineSession):
training_request.pop("training_job_name", None)
# Keep training_job_name in the request. The TrainingStep strips it via
# trim_request_dict, which drops it by default but preserves the base_job_name
# prefix when PipelineDefinitionConfig(use_custom_job_prefix=True). Popping it here
# unconditionally left use_custom_job_prefix nothing to preserve (issue #5776, #6299).
# Convert snake_case to PascalCase for AWS API
pipeline_request = {to_pascal_case(k): v for k, v in training_request.items()}
serialized_request = serialize(pipeline_request)
Expand Down
6 changes: 4 additions & 2 deletions sagemaker-train/src/sagemaker/train/tuner.py
Original file line number Diff line number Diff line change
Expand Up @@ -1290,8 +1290,10 @@ def _start_tuning_job(self, inputs):
from sagemaker.core.utils.utils import serialize
from sagemaker.core.apiutils._boto_functions import to_pascal_case

# Remove job name for pipeline as it's auto-generated at execution time
tuning_request.pop("hyper_parameter_tuning_job_name", None)
# Keep hyper_parameter_tuning_job_name in the request. The TuningStep strips it via
# trim_request_dict, which drops it by default but preserves the base_tuning_job_name
# prefix when PipelineDefinitionConfig(use_custom_job_prefix=True). Popping it here
# unconditionally left use_custom_job_prefix nothing to preserve (issue #6299, #5776).
# Convert snake_case to PascalCase for AWS API
pipeline_request = {to_pascal_case(k): v for k, v in tuning_request.items()}
serialized_request = serialize(pipeline_request)
Expand Down
64 changes: 59 additions & 5 deletions sagemaker-train/tests/unit/train/test_model_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -254,6 +254,60 @@ def test_model_trainer_param_validation(test_case, modules_session):
assert trainer.base_job_name == DEFAULT_BASE_NAME


def test_pipeline_session_request_keeps_training_job_name(modules_session):
"""Regression for #5776 / #6299.

Under a PipelineSession the request must KEEP TrainingJobName. The TrainingStep strips it via
trim_request_dict (dropped by default, prefix preserved when use_custom_job_prefix=True);
popping it here left use_custom_job_prefix nothing to preserve.
"""
from unittest.mock import Mock
from sagemaker.core.workflow.pipeline_context import PipelineSession

session = Mock(spec=PipelineSession)
session.default_bucket.return_value = DEFAULT_BUCKET
session.default_bucket_prefix = DEFAULT_BUCKET_PREFIX
session.boto_region_name = DEFAULT_REGION
session.sagemaker_config = {}

trainer = ModelTrainer(
training_image=DEFAULT_IMAGE,
role=DEFAULT_ROLE,
compute=DEFAULT_COMPUTE_CONFIG,
stopping_condition=DEFAULT_STOPPING_CONDITION,
output_data_config=DEFAULT_OUTPUT_DATA_CONFIG,
base_job_name="my-prefix",
sagemaker_session=session,
)

args = trainer._create_training_job_args(input_data_config=[])

assert "TrainingJobName" in args, "TrainingJobName must survive into the pipeline request"
assert args["TrainingJobName"].startswith("my-prefix")

# The generated name must use a timestamp format base_from_name can strip, otherwise
# trim_request_dict bakes a definition-time timestamp into the pipeline definition
# instead of preserving the base prefix.
from sagemaker.core.common_utils import base_from_name

assert base_from_name(args["TrainingJobName"]) == "my-prefix"

# End-to-end: with use_custom_job_prefix=True, trim_request_dict must leave exactly the base.
from types import SimpleNamespace
from sagemaker.core.workflow.utilities import trim_request_dict
from sagemaker.core.workflow.pipeline_definition_config import PipelineDefinitionConfig

custom_prefix_config = SimpleNamespace(
pipeline_definition_config=PipelineDefinitionConfig(use_custom_job_prefix=True)
)
trimmed = trim_request_dict(dict(args), "TrainingJobName", custom_prefix_config)
assert trimmed["TrainingJobName"] == "my-prefix"

# And by default (no config) the key is dropped entirely.
dropped = trim_request_dict(dict(args), "TrainingJobName", None)
assert "TrainingJobName" not in dropped


@patch("sagemaker.train.model_trainer.TrainingJob")
def test_train_with_default_params(mock_training_job, model_trainer):
model_trainer.train()
Expand Down Expand Up @@ -871,7 +925,7 @@ def test_remote_debug_config(mock_training_job, modules_session):
)


@patch("sagemaker.train.model_trainer._get_unique_name")
@patch("sagemaker.train.model_trainer.name_from_base")
@patch("sagemaker.train.model_trainer.TrainingJob")
def test_model_trainer_full_init(mock_training_job, mock_unique_name, modules_session):
def mock_upload_data(path, bucket, key_prefix):
Expand Down Expand Up @@ -1147,7 +1201,7 @@ def mock_upload_data(path, bucket, key_prefix):

# TODO: Re-Enable after local mode fully migrated to v3
# @patch("sagemaker.train.model_trainer._LocalContainer")
# @patch("sagemaker.train.model_trainer._get_unique_name")
# @patch("sagemaker.train.model_trainer.name_from_base")
# @patch("sagemaker.train.local.local_container.download_folder")
# def test_model_trainer_local_full_init(
# mock_download_folder, mock_unique_name, mock_local_container, modules_session
Expand Down Expand Up @@ -1403,7 +1457,7 @@ def test_hyperparameters_invalid(mock_exists, modules_session):
)


@patch("sagemaker.train.model_trainer._get_unique_name")
@patch("sagemaker.train.model_trainer.name_from_base")
@patch("sagemaker.train.model_trainer.TrainingJob")
def test_model_trainer_default_paths(mock_training_job, mock_unique_name, modules_session):
def mock_upload_data(path, bucket, key_prefix):
Expand Down Expand Up @@ -1507,7 +1561,7 @@ def test_metric_definitions(mock_training_job, modules_session):
)


@patch("sagemaker.train.model_trainer._get_unique_name")
@patch("sagemaker.train.model_trainer.name_from_base")
@patch("sagemaker.core.resources.TrainingJob")
def test_nova_recipe(mock_training_job, mock_unique_name, modules_session):
def mock_upload_data(path, bucket, key_prefix):
Expand Down Expand Up @@ -1779,7 +1833,7 @@ def test_nova_recipe_model_package_config_only_mpg_from_recipe(modules_session):
os.unlink(recipe.name)


@patch("sagemaker.train.model_trainer._get_unique_name")
@patch("sagemaker.train.model_trainer.name_from_base")
@patch("sagemaker.train.model_trainer.TrainingJob")
def test_llmft_recipe(mock_training_job, mock_unique_name, modules_session):
def mock_upload_data(path, bucket, key_prefix):
Expand Down
34 changes: 34 additions & 0 deletions sagemaker-train/tests/unit/train/test_tuner.py
Original file line number Diff line number Diff line change
Expand Up @@ -598,6 +598,40 @@ def test_build_training_job_definition_includes_spot_params(self):
definition.stopping_condition.max_wait_time_in_seconds, int
), "Max wait time should be set"

def test_pipeline_session_request_keeps_tuning_job_name(self):
"""Regression for #6299 / #5776.

Under a PipelineSession the tuning request must KEEP HyperParameterTuningJobName. The
TuningStep strips it via trim_request_dict (dropped by default, prefix preserved when
use_custom_job_prefix=True); popping it here left use_custom_job_prefix nothing to preserve.
"""
from unittest.mock import Mock
from sagemaker.core.workflow.pipeline_context import PipelineSession

pipeline_session = Mock(spec=PipelineSession)
mock_trainer = _create_mock_model_trainer()
mock_trainer.sagemaker_session = pipeline_session

tuner = HyperparameterTuner(
model_trainer=mock_trainer,
objective_metric_name="accuracy",
hyperparameter_ranges=_create_single_hp_range(),
)
tuner._current_job_name = "my-prefix-2026-01-01-00-00-00-000"

tuner._start_tuning_job(inputs=None)

pipeline_session._intercept_create_request.assert_called_once()
serialized_request = pipeline_session._intercept_create_request.call_args.args[0]
assert "HyperParameterTuningJobName" in serialized_request
assert serialized_request["HyperParameterTuningJobName"].startswith("my-prefix")

# The name must be strippable back to the base, or trim_request_dict would preserve a
# definition-time timestamp instead of the prefix when use_custom_job_prefix=True.
from sagemaker.core.common_utils import base_from_name

assert base_from_name(serialized_request["HyperParameterTuningJobName"]) == "my-prefix"

def test_build_training_job_definition_includes_environment_variables(self):
"""Test that _build_training_job_definition includes environment variables.

Expand Down