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: 16 additions & 0 deletions sagemaker-train/src/sagemaker/train/common.py
Original file line number Diff line number Diff line change
Expand Up @@ -84,6 +84,22 @@ def to_dict(self) -> Dict[str, Any]:
def to_user_dict(self) -> Dict[str, Any]:
"""Return only user-explicitly-set hyperparameters as string key-value pairs."""
return {k: str(getattr(self, k)) for k in self._user_set if getattr(self, k, None) is not None}

def required_keys(self) -> set:
"""Return the set of spec keys marked ``required``.

These are hyperparameters the recipe/model requires a value for. They
must survive into the final training request; ``to_dict()`` skips any
spec whose value is ``None``, so a required parameter with no default
that the user never set would otherwise be dropped silently. Callers
use this set to surface such omissions instead (see
``_validate_hyperparameter_values``).
"""
return {
name
for name, spec in self._specs.items()
if isinstance(spec, dict) and spec.get("required")
}

def __setattr__(self, name: str, value: Any):
if name.startswith('_'):
Expand Down
40 changes: 38 additions & 2 deletions sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1338,8 +1338,44 @@ def _validate_s3_path_exists(s3_path: str, sagemaker_session):
raise ValueError(f"Failed to validate/create S3 path '{s3_path}': {str(e)}")


def _validate_hyperparameter_values(hyperparameters: dict):
"""Validate hyperparameter values for allowed characters."""
def _validate_hyperparameter_values(hyperparameters: dict, options: Optional["FineTuningOptions"] = None):
"""Validate hyperparameter values for allowed characters.

When ``options`` (the trainer's ``FineTuningOptions``) is provided, this
also surfaces required hyperparameters that are missing from the final
request. ``FineTuningOptions.to_dict()`` silently skips any spec whose
value is ``None``, so a required parameter with no default that the user
never set (and that no recipe/override supplied) would otherwise be dropped
without any error or warning, and the training job would launch
mis-configured. Raising here fails fast, client-side, with an actionable
message instead.

Args:
hyperparameters: The final, fully merged hyperparameters dict that will
be sent to the training job.
options: Optional ``FineTuningOptions`` describing the spec. Only passed
from call sites that run *after* recipe/override merge, so a
required value supplied by the recipe is correctly counted as
present.
"""
# Surface (don't silently drop) required hyperparameters missing from the
# final request. Guarded on the concrete type so mocks / other objects are
# ignored.
if isinstance(options, FineTuningOptions):
missing = sorted(
key
for key in options.required_keys()
if hyperparameters.get(key) in (None, "")
)
if missing:
raise ValueError(
"Missing required hyperparameter(s): "
f"{', '.join(missing)}. Set them via "
"`trainer.hyperparameters.<name> = <value>` (or supply them in a "
"recipe / overrides) before training. These parameters are "
"required and cannot be omitted from the training request."
)

import re
allowed_chars = r"^[a-zA-Z0-9/_.:,\-\s'\"\[\]]*$"
for key, value in hyperparameters.items():
Expand Down
2 changes: 1 addition & 1 deletion sagemaker-train/src/sagemaker/train/dpo_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -339,7 +339,7 @@ def train(self,
if effective_training_dataset is not None:
self.is_multimodal = is_multimodal_data(effective_training_dataset)

_validate_hyperparameter_values(final_hyperparameters)
_validate_hyperparameter_values(final_hyperparameters, self.hyperparameters)

model_package_config = _create_model_package_config(
model_package_group_name=self.model_package_group,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -309,7 +309,7 @@ def train(
# Apply recipe/overrides if provided (overrides > recipe > Hub defaults)
self._final_hyperparameters = self._apply_recipe_to_hyperparameters(self._final_hyperparameters)

_validate_hyperparameter_values(self._final_hyperparameters)
_validate_hyperparameter_values(self._final_hyperparameters, self.hyperparameters)

if training_dataset is not None:
self.training_dataset = training_dataset
Expand Down
2 changes: 1 addition & 1 deletion sagemaker-train/src/sagemaker/train/rlaif_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -305,7 +305,7 @@ def train(self, training_dataset: Optional[Union[str, DataSet]] = None, validati
if effective_training_dataset is not None:
self.is_multimodal = is_multimodal_data(effective_training_dataset)

_validate_hyperparameter_values(final_hyperparameters)
_validate_hyperparameter_values(final_hyperparameters, self.hyperparameters)

model_package_config = _create_model_package_config(
model_package_group_name=self.model_package_group,
Expand Down
2 changes: 1 addition & 1 deletion sagemaker-train/src/sagemaker/train/rlvr_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -518,7 +518,7 @@ def train(self, training_dataset: Optional[Union[str, DataSet]] = None,
self.is_multimodal = is_multimodal_data(effective_training_dataset)

# Validate hyperparameter values
_validate_hyperparameter_values(final_hyperparameters)
_validate_hyperparameter_values(final_hyperparameters, self.hyperparameters)

model_package_config = _create_model_package_config(
model_package_group_name=self.model_package_group,
Expand Down
2 changes: 1 addition & 1 deletion sagemaker-train/src/sagemaker/train/sft_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -407,7 +407,7 @@ def train(self, training_dataset: Optional[Union[str, DataSet]] = None, validati
final_hyperparameters[param_name] = param_value

# Validate hyperparameter values
_validate_hyperparameter_values(final_hyperparameters)
_validate_hyperparameter_values(final_hyperparameters, self.hyperparameters)

model_package_config = _create_model_package_config(
model_package_group_name=self.model_package_group,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1784,3 +1784,93 @@ def test_list_hyperparameters_accepts_enum_values(self, mock_boto_client, mock_g
)

assert result.learning_rate == 0.0001


class TestValidateHyperparameterValues:
"""Tests for _validate_hyperparameter_values, including the required-key
surfacing that prevents required hyperparameters from being silently
dropped from the training request.
"""

def test_missing_required_hyperparameter_raises(self):
"""A required spec with no value must be surfaced, not silently dropped."""
from sagemaker.train.common import FineTuningOptions

options = FineTuningOptions({
"learning_rate": {"type": "float", "default": 0.0001},
"required_no_default": {"type": "string", "required": True},
})
# to_dict() drops required_no_default (its value is None).
final = options.to_dict()
assert "required_no_default" not in final

with pytest.raises(ValueError, match="Missing required hyperparameter"):
fu._validate_hyperparameter_values(final, options)

def test_missing_required_error_names_the_keys(self):
from sagemaker.train.common import FineTuningOptions

options = FineTuningOptions({
"req_a": {"type": "string", "required": True},
"req_b": {"type": "string", "required": True},
})
with pytest.raises(ValueError) as exc:
fu._validate_hyperparameter_values(options.to_dict(), options)
msg = str(exc.value)
assert "req_a" in msg and "req_b" in msg

def test_required_present_passes(self):
"""When the required value is set, validation passes."""
from sagemaker.train.common import FineTuningOptions

options = FineTuningOptions({
"required_no_default": {"type": "string", "required": True},
})
options.required_no_default = "some-value"
# No raise.
fu._validate_hyperparameter_values(options.to_dict(), options)

def test_required_supplied_by_recipe_merge_passes(self):
"""A required key absent from the FineTuningOptions object but present in
the final merged dict (e.g. supplied by a recipe/override) passes."""
from sagemaker.train.common import FineTuningOptions

options = FineTuningOptions({
"required_no_default": {"type": "string", "required": True},
})
# Simulate recipe/override merge populating the final request dict.
final = {"required_no_default": "from-recipe"}
fu._validate_hyperparameter_values(final, options) # no raise

def test_required_empty_string_is_treated_as_missing(self):
from sagemaker.train.common import FineTuningOptions

options = FineTuningOptions({
"required_no_default": {"type": "string", "required": True},
})
with pytest.raises(ValueError, match="Missing required hyperparameter"):
fu._validate_hyperparameter_values({"required_no_default": ""}, options)

def test_no_options_is_backward_compatible(self):
"""Called without options (e.g. the pre-merge base_trainer call site),
only character validation runs — no required check."""
# Would raise if the required check ran, but no options are passed.
fu._validate_hyperparameter_values({"learning_rate": "0.1"}) # no raise

def test_non_finetuning_options_is_ignored(self):
"""A Mock (or any non-FineTuningOptions) passed as options is ignored,
so existing mock-based trainer tests keep working."""
fu._validate_hyperparameter_values({"learning_rate": "0.1"}, Mock()) # no raise

def test_invalid_characters_still_raise(self):
"""The original character validation is preserved."""
with pytest.raises(ValueError, match="invalid characters"):
fu._validate_hyperparameter_values({"bad": "value;with;semicolons"})

def test_no_required_keys_passes(self):
from sagemaker.train.common import FineTuningOptions

options = FineTuningOptions({
"learning_rate": {"type": "float", "default": 0.0001},
})
fu._validate_hyperparameter_values(options.to_dict(), options) # no raise
34 changes: 34 additions & 0 deletions sagemaker-train/tests/unit/train/test_common.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,40 @@ def test_to_dict_all_none_returns_empty(self):
assert result == {}


class TestFineTuningOptionsRequiredKeys:
"""Tests for FineTuningOptions.required_keys()."""

def test_returns_only_required_specs(self):
options = FineTuningOptions({
"learning_rate": {"default": 0.001, "type": "float"},
"required_no_default": {"type": "string", "required": True},
"required_with_default": {"default": 5, "type": "integer", "required": True},
})
assert options.required_keys() == {"required_no_default", "required_with_default"}

def test_returns_empty_when_none_required(self):
options = FineTuningOptions({
"learning_rate": {"default": 0.001, "type": "float"},
"epochs": {"default": 3, "type": "integer"},
})
assert options.required_keys() == set()

def test_required_false_is_excluded(self):
options = FineTuningOptions({
"opt_in": {"default": None, "type": "string", "required": False},
})
assert options.required_keys() == set()

def test_required_key_dropped_by_to_dict_is_still_reported(self):
"""A required spec with no default (None) is dropped by to_dict() but
must still be reported by required_keys() so callers can surface it."""
options = FineTuningOptions({
"required_no_default": {"type": "string", "required": True},
})
assert "required_no_default" not in options.to_dict()
assert options.required_keys() == {"required_no_default"}


import pytest


Expand Down
Loading