Skip to content
Merged
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
40 changes: 17 additions & 23 deletions sagemaker-train/src/sagemaker/train/base_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -176,41 +176,35 @@ def _apply_user_hyperparameters(self, user_hyperparameters: Optional[Dict[str, A
"""Apply constructor-supplied hyperparameters onto the resolved FineTuningOptions.

The fine-tuning trainers replace ``self.hyperparameters`` with a
``FineTuningOptions`` built from the model's Hub spec, which would otherwise
discard any ``hyperparameters`` dict passed at construction. This re-applies
those user-provided values, but only for names that are overridable for the
model (i.e. present in the options' ``_specs``). Each applied value goes through
``FineTuningOptions.__setattr__``, so it is still validated against the spec and
an out-of-spec value for an overridable name raises, exactly as
``trainer.hyperparameters.<name> = value`` would.
``FineTuningOptions`` built from the model's recipe override-params spec, which
would otherwise discard any ``hyperparameters`` dict passed at construction. This
re-applies those user-provided values by routing each one through
``FineTuningOptions.__setattr__``, so a dict passed at construction behaves
exactly like the ``trainer.hyperparameters.<name> = value`` path and is validated
against the recipe spec:

Names that are not overridable are ignored (not applied), and a single warning
lists them so the user knows those values will not take effect.
* an unknown option name raises ``AttributeError``;
* an out-of-spec or off-enum value for a known name raises ``ValueError``.

Values are never silently dropped in favor of the recipe default, which would
otherwise launch a billable job on a configuration the caller did not set.

No-op when nothing was supplied or when ``self.hyperparameters`` is not a
spec-backed ``FineTuningOptions`` (e.g. a plain dict).
spec-backed ``FineTuningOptions`` (e.g. ``ModelTrainer``'s plain dict).

Args:
user_hyperparameters: The hyperparameters dict captured from construction.

Raises:
AttributeError: If a supplied name is not a valid option for the recipe.
ValueError: If a supplied value fails the recipe spec (type/range/enum).
"""
if not user_hyperparameters:
return
specs = getattr(getattr(self, "hyperparameters", None), "_specs", None)
if not isinstance(specs, dict):
if not isinstance(getattr(getattr(self, "hyperparameters", None), "_specs", None), dict):
return
ignored = []
for name, value in user_hyperparameters.items():
if name not in specs:
ignored.append(name)
continue
setattr(self.hyperparameters, name, value)
if ignored:
logger.warning(
"Ignoring hyperparameters that are not overridable for this model: %s. "
"These values will not take effect. Overridable hyperparameters: %s",
ignored,
list(specs.keys()),
)

def _is_nova_model_for_telemetry(self) -> bool:
"""Check if the model is a Nova model for telemetry tracking."""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -41,18 +41,12 @@ def test_applies_valid_values_and_marks_user_set():
assert trainer.hyperparameters._user_set == {"learning_rate", "epochs"}


def test_invalid_option_name_is_ignored_with_warning(caplog):
def test_invalid_option_name_raises():
options = _make_options()
with caplog.at_level("WARNING"):
trainer = _apply(options, {"not_a_real_option": 1, "learning_rate": 0.001})

# The overridable value is applied; the non-overridable one is skipped.
assert trainer.hyperparameters.learning_rate == 0.001
assert not hasattr(trainer.hyperparameters, "not_a_real_option")
assert trainer.hyperparameters._user_set == {"learning_rate"}
# A warning names the ignored, non-overridable hyperparameter.
assert "not_a_real_option" in caplog.text
assert "not overridable" in caplog.text.lower()
# An unknown option name raises (consistent with trainer.hyperparameters.x = v),
# rather than being silently ignored.
with pytest.raises(AttributeError):
_apply(options, {"not_a_real_option": 1})


def test_out_of_spec_value_raises():
Expand Down
18 changes: 18 additions & 0 deletions sagemaker-train/tests/unit/train/test_dpo_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,24 @@ def test_init_with_defaults(self, mock_finetuning_options, mock_validate_group,
assert trainer.training_type == TrainingType.LORA
assert trainer.model == "test-model"

@patch('sagemaker.train.dpo_trainer._validate_and_resolve_model_package_group')
@patch('sagemaker.train.dpo_trainer._get_fine_tuning_options_and_model_arn')
def test_init_applies_constructor_hyperparameters(self, mock_finetuning_options, mock_validate_group, mock_session):
"""Constructor hyperparameters are applied onto the resolved FineTuningOptions (wiring guard)."""
from sagemaker.train.common import FineTuningOptions
mock_validate_group.return_value = "test-group"
options = FineTuningOptions(
{"learning_rate": {"type": "float", "default": 0.0001, "min": 0.0, "max": 1.0}}
)
mock_finetuning_options.return_value = (options, "model-arn", False)
trainer = DPOTrainer(
model="test-model",
model_package_group="test-group",
hyperparameters={"learning_rate": 0.001},
)
assert trainer.hyperparameters.learning_rate == 0.001
assert "learning_rate" in trainer.hyperparameters._user_set

@patch('sagemaker.train.dpo_trainer._validate_and_resolve_model_package_group')
@patch('sagemaker.train.dpo_trainer._get_fine_tuning_options_and_model_arn')
def test_init_with_full_training_type(self, mock_finetuning_options, mock_validate_group, mock_session):
Expand Down
18 changes: 18 additions & 0 deletions sagemaker-train/tests/unit/train/test_rlaif_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,24 @@ def test_init_with_defaults(self, mock_finetuning_options, mock_validate_group,
assert trainer.training_type == TrainingType.LORA
assert trainer.model == "test-model"

@patch('sagemaker.train.rlaif_trainer._validate_and_resolve_model_package_group')
@patch('sagemaker.train.rlaif_trainer._get_fine_tuning_options_and_model_arn')
def test_init_applies_constructor_hyperparameters(self, mock_finetuning_options, mock_validate_group, mock_session):
"""Constructor hyperparameters are applied onto the resolved FineTuningOptions (wiring guard)."""
from sagemaker.train.common import FineTuningOptions
mock_validate_group.return_value = "test-group"
options = FineTuningOptions(
{"learning_rate": {"type": "float", "default": 0.0001, "min": 0.0, "max": 1.0}}
)
mock_finetuning_options.return_value = (options, "model-arn", False)
trainer = RLAIFTrainer(
model="test-model",
model_package_group="test-group",
hyperparameters={"learning_rate": 0.001},
)
assert trainer.hyperparameters.learning_rate == 0.001
assert "learning_rate" in trainer.hyperparameters._user_set

@patch('sagemaker.train.rlaif_trainer._validate_and_resolve_model_package_group')
@patch('sagemaker.train.rlaif_trainer._get_fine_tuning_options_and_model_arn')
def test_init_with_full_training_type(self, mock_finetuning_options, mock_validate_group, mock_session):
Expand Down
18 changes: 18 additions & 0 deletions sagemaker-train/tests/unit/train/test_rlvr_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,24 @@ def test_init_with_defaults(self, mock_finetuning_options, mock_validate_group,
assert trainer.training_type == TrainingType.LORA
assert trainer.model == "test-model"

@patch('sagemaker.train.rlvr_trainer._validate_and_resolve_model_package_group')
@patch('sagemaker.train.rlvr_trainer._get_fine_tuning_options_and_model_arn')
def test_init_applies_constructor_hyperparameters(self, mock_finetuning_options, mock_validate_group, mock_session):
"""Constructor hyperparameters are applied onto the resolved FineTuningOptions (wiring guard)."""
from sagemaker.train.common import FineTuningOptions
mock_validate_group.return_value = "test-group"
options = FineTuningOptions(
{"learning_rate": {"type": "float", "default": 0.0001, "min": 0.0, "max": 1.0}}
)
mock_finetuning_options.return_value = (options, "model-arn", False)
trainer = RLVRTrainer(
model="test-model",
model_package_group="test-group",
hyperparameters={"learning_rate": 0.001},
)
assert trainer.hyperparameters.learning_rate == 0.001
assert "learning_rate" in trainer.hyperparameters._user_set

@patch('sagemaker.train.rlvr_trainer._validate_and_resolve_model_package_group')
@patch('sagemaker.train.rlvr_trainer._get_fine_tuning_options_and_model_arn')
def test_init_with_full_training_type(self, mock_finetuning_options, mock_validate_group, mock_session):
Expand Down
19 changes: 8 additions & 11 deletions sagemaker-train/tests/unit/train/test_sft_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,24 +47,21 @@ def test_init_applies_constructor_hyperparameters(self, mock_finetuning_options,

@patch('sagemaker.train.sft_trainer._validate_and_resolve_model_package_group')
@patch('sagemaker.train.sft_trainer._get_fine_tuning_options_and_model_arn')
def test_init_ignores_non_overridable_constructor_hyperparameter(self, mock_finetuning_options, mock_validate_group, mock_session):
"""A non-overridable constructor hyperparameter is ignored (not applied, no raise)."""
def test_init_invalid_constructor_hyperparameter_raises(self, mock_finetuning_options, mock_validate_group, mock_session):
"""A non-overridable constructor hyperparameter raises rather than being dropped."""
from sagemaker.train.common import FineTuningOptions
mock_validate_group.return_value = "test-group"
options = FineTuningOptions(
{"learning_rate": {"type": "float", "default": 0.0001, "min": 0.0, "max": 1.0}}
)
mock_finetuning_options.return_value = (options, "model-arn", False)

trainer = SFTTrainer(
model="test-model",
model_package_group="test-group",
hyperparameters={"learning_rate": 0.001, "not_a_real_option": 1},
)

# Overridable value applied; non-overridable one ignored rather than raising.
assert trainer.hyperparameters.learning_rate == 0.001
assert not hasattr(trainer.hyperparameters, "not_a_real_option")
with pytest.raises(AttributeError):
SFTTrainer(
model="test-model",
model_package_group="test-group",
hyperparameters={"learning_rate": 0.001, "not_a_real_option": 1},
)

@patch('sagemaker.train.sft_trainer._validate_and_resolve_model_package_group')
@patch('sagemaker.train.sft_trainer._get_fine_tuning_options_and_model_arn')
Expand Down
Loading