Skip to content

Commit 0361ea7

Browse files
author
Roja Reddy Sareddy
committed
fix(train): surface required hyperparameters instead of dropping them
FineTuningOptions.to_dict() builds the hyperparameters dict from self._specs but omits any None value. A spec marked required=True with no default that the user never set was therefore silently dropped from the training request, so a job could launch missing a required hyperparameter and only surface the problem via bad/failed results. Add FineTuningOptions.required_keys() and have the shared _validate_hyperparameter_values() raise a clear error listing any required hyperparameter missing from the final (post recipe/override merge) request. The check is opt-in via an options arg passed only at the post-merge leaf-trainer call sites (SFT/DPO/RLVR/RLAIF/MultiTurnRL); the pre-merge base_trainer call is left unchanged to avoid false positives. Guarded on the concrete type so mocks are ignored. Adds unit tests for required_keys() and the surfacing behavior.
1 parent 8d924cb commit 0361ea7

9 files changed

Lines changed: 183 additions & 7 deletions

File tree

sagemaker-train/src/sagemaker/train/common.py

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -84,6 +84,22 @@ def to_dict(self) -> Dict[str, Any]:
8484
def to_user_dict(self) -> Dict[str, Any]:
8585
"""Return only user-explicitly-set hyperparameters as string key-value pairs."""
8686
return {k: str(getattr(self, k)) for k in self._user_set if getattr(self, k, None) is not None}
87+
88+
def required_keys(self) -> set:
89+
"""Return the set of spec keys marked ``required``.
90+
91+
These are hyperparameters the recipe/model requires a value for. They
92+
must survive into the final training request; ``to_dict()`` skips any
93+
spec whose value is ``None``, so a required parameter with no default
94+
that the user never set would otherwise be dropped silently. Callers
95+
use this set to surface such omissions instead (see
96+
``_validate_hyperparameter_values``).
97+
"""
98+
return {
99+
name
100+
for name, spec in self._specs.items()
101+
if isinstance(spec, dict) and spec.get("required")
102+
}
87103

88104
def __setattr__(self, name: str, value: Any):
89105
if name.startswith('_'):

sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py

Lines changed: 38 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1338,8 +1338,44 @@ def _validate_s3_path_exists(s3_path: str, sagemaker_session):
13381338
raise ValueError(f"Failed to validate/create S3 path '{s3_path}': {str(e)}")
13391339

13401340

1341-
def _validate_hyperparameter_values(hyperparameters: dict):
1342-
"""Validate hyperparameter values for allowed characters."""
1341+
def _validate_hyperparameter_values(hyperparameters: dict, options: Optional["FineTuningOptions"] = None):
1342+
"""Validate hyperparameter values for allowed characters.
1343+
1344+
When ``options`` (the trainer's ``FineTuningOptions``) is provided, this
1345+
also surfaces required hyperparameters that are missing from the final
1346+
request. ``FineTuningOptions.to_dict()`` silently skips any spec whose
1347+
value is ``None``, so a required parameter with no default that the user
1348+
never set (and that no recipe/override supplied) would otherwise be dropped
1349+
without any error or warning, and the training job would launch
1350+
mis-configured. Raising here fails fast, client-side, with an actionable
1351+
message instead.
1352+
1353+
Args:
1354+
hyperparameters: The final, fully merged hyperparameters dict that will
1355+
be sent to the training job.
1356+
options: Optional ``FineTuningOptions`` describing the spec. Only passed
1357+
from call sites that run *after* recipe/override merge, so a
1358+
required value supplied by the recipe is correctly counted as
1359+
present.
1360+
"""
1361+
# Surface (don't silently drop) required hyperparameters missing from the
1362+
# final request. Guarded on the concrete type so mocks / other objects are
1363+
# ignored.
1364+
if isinstance(options, FineTuningOptions):
1365+
missing = sorted(
1366+
key
1367+
for key in options.required_keys()
1368+
if hyperparameters.get(key) in (None, "")
1369+
)
1370+
if missing:
1371+
raise ValueError(
1372+
"Missing required hyperparameter(s): "
1373+
f"{', '.join(missing)}. Set them via "
1374+
"`trainer.hyperparameters.<name> = <value>` (or supply them in a "
1375+
"recipe / overrides) before training. These parameters are "
1376+
"required and cannot be omitted from the training request."
1377+
)
1378+
13431379
import re
13441380
allowed_chars = r"^[a-zA-Z0-9/_.:,\-\s'\"\[\]]*$"
13451381
for key, value in hyperparameters.items():

sagemaker-train/src/sagemaker/train/dpo_trainer.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -339,7 +339,7 @@ def train(self,
339339
if effective_training_dataset is not None:
340340
self.is_multimodal = is_multimodal_data(effective_training_dataset)
341341

342-
_validate_hyperparameter_values(final_hyperparameters)
342+
_validate_hyperparameter_values(final_hyperparameters, self.hyperparameters)
343343

344344
model_package_config = _create_model_package_config(
345345
model_package_group_name=self.model_package_group,

sagemaker-train/src/sagemaker/train/multi_turn_rl_trainer.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -309,7 +309,7 @@ def train(
309309
# Apply recipe/overrides if provided (overrides > recipe > Hub defaults)
310310
self._final_hyperparameters = self._apply_recipe_to_hyperparameters(self._final_hyperparameters)
311311

312-
_validate_hyperparameter_values(self._final_hyperparameters)
312+
_validate_hyperparameter_values(self._final_hyperparameters, self.hyperparameters)
313313

314314
if training_dataset is not None:
315315
self.training_dataset = training_dataset

sagemaker-train/src/sagemaker/train/rlaif_trainer.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -305,7 +305,7 @@ def train(self, training_dataset: Optional[Union[str, DataSet]] = None, validati
305305
if effective_training_dataset is not None:
306306
self.is_multimodal = is_multimodal_data(effective_training_dataset)
307307

308-
_validate_hyperparameter_values(final_hyperparameters)
308+
_validate_hyperparameter_values(final_hyperparameters, self.hyperparameters)
309309

310310
model_package_config = _create_model_package_config(
311311
model_package_group_name=self.model_package_group,

sagemaker-train/src/sagemaker/train/rlvr_trainer.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -518,7 +518,7 @@ def train(self, training_dataset: Optional[Union[str, DataSet]] = None,
518518
self.is_multimodal = is_multimodal_data(effective_training_dataset)
519519

520520
# Validate hyperparameter values
521-
_validate_hyperparameter_values(final_hyperparameters)
521+
_validate_hyperparameter_values(final_hyperparameters, self.hyperparameters)
522522

523523
model_package_config = _create_model_package_config(
524524
model_package_group_name=self.model_package_group,

sagemaker-train/src/sagemaker/train/sft_trainer.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -407,7 +407,7 @@ def train(self, training_dataset: Optional[Union[str, DataSet]] = None, validati
407407
final_hyperparameters[param_name] = param_value
408408

409409
# Validate hyperparameter values
410-
_validate_hyperparameter_values(final_hyperparameters)
410+
_validate_hyperparameter_values(final_hyperparameters, self.hyperparameters)
411411

412412
model_package_config = _create_model_package_config(
413413
model_package_group_name=self.model_package_group,

sagemaker-train/tests/unit/train/common_utils/test_finetune_utils.py

Lines changed: 90 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1784,3 +1784,93 @@ def test_list_hyperparameters_accepts_enum_values(self, mock_boto_client, mock_g
17841784
)
17851785

17861786
assert result.learning_rate == 0.0001
1787+
1788+
1789+
class TestValidateHyperparameterValues:
1790+
"""Tests for _validate_hyperparameter_values, including the required-key
1791+
surfacing that prevents required hyperparameters from being silently
1792+
dropped from the training request.
1793+
"""
1794+
1795+
def test_missing_required_hyperparameter_raises(self):
1796+
"""A required spec with no value must be surfaced, not silently dropped."""
1797+
from sagemaker.train.common import FineTuningOptions
1798+
1799+
options = FineTuningOptions({
1800+
"learning_rate": {"type": "float", "default": 0.0001},
1801+
"required_no_default": {"type": "string", "required": True},
1802+
})
1803+
# to_dict() drops required_no_default (its value is None).
1804+
final = options.to_dict()
1805+
assert "required_no_default" not in final
1806+
1807+
with pytest.raises(ValueError, match="Missing required hyperparameter"):
1808+
fu._validate_hyperparameter_values(final, options)
1809+
1810+
def test_missing_required_error_names_the_keys(self):
1811+
from sagemaker.train.common import FineTuningOptions
1812+
1813+
options = FineTuningOptions({
1814+
"req_a": {"type": "string", "required": True},
1815+
"req_b": {"type": "string", "required": True},
1816+
})
1817+
with pytest.raises(ValueError) as exc:
1818+
fu._validate_hyperparameter_values(options.to_dict(), options)
1819+
msg = str(exc.value)
1820+
assert "req_a" in msg and "req_b" in msg
1821+
1822+
def test_required_present_passes(self):
1823+
"""When the required value is set, validation passes."""
1824+
from sagemaker.train.common import FineTuningOptions
1825+
1826+
options = FineTuningOptions({
1827+
"required_no_default": {"type": "string", "required": True},
1828+
})
1829+
options.required_no_default = "some-value"
1830+
# No raise.
1831+
fu._validate_hyperparameter_values(options.to_dict(), options)
1832+
1833+
def test_required_supplied_by_recipe_merge_passes(self):
1834+
"""A required key absent from the FineTuningOptions object but present in
1835+
the final merged dict (e.g. supplied by a recipe/override) passes."""
1836+
from sagemaker.train.common import FineTuningOptions
1837+
1838+
options = FineTuningOptions({
1839+
"required_no_default": {"type": "string", "required": True},
1840+
})
1841+
# Simulate recipe/override merge populating the final request dict.
1842+
final = {"required_no_default": "from-recipe"}
1843+
fu._validate_hyperparameter_values(final, options) # no raise
1844+
1845+
def test_required_empty_string_is_treated_as_missing(self):
1846+
from sagemaker.train.common import FineTuningOptions
1847+
1848+
options = FineTuningOptions({
1849+
"required_no_default": {"type": "string", "required": True},
1850+
})
1851+
with pytest.raises(ValueError, match="Missing required hyperparameter"):
1852+
fu._validate_hyperparameter_values({"required_no_default": ""}, options)
1853+
1854+
def test_no_options_is_backward_compatible(self):
1855+
"""Called without options (e.g. the pre-merge base_trainer call site),
1856+
only character validation runs — no required check."""
1857+
# Would raise if the required check ran, but no options are passed.
1858+
fu._validate_hyperparameter_values({"learning_rate": "0.1"}) # no raise
1859+
1860+
def test_non_finetuning_options_is_ignored(self):
1861+
"""A Mock (or any non-FineTuningOptions) passed as options is ignored,
1862+
so existing mock-based trainer tests keep working."""
1863+
fu._validate_hyperparameter_values({"learning_rate": "0.1"}, Mock()) # no raise
1864+
1865+
def test_invalid_characters_still_raise(self):
1866+
"""The original character validation is preserved."""
1867+
with pytest.raises(ValueError, match="invalid characters"):
1868+
fu._validate_hyperparameter_values({"bad": "value;with;semicolons"})
1869+
1870+
def test_no_required_keys_passes(self):
1871+
from sagemaker.train.common import FineTuningOptions
1872+
1873+
options = FineTuningOptions({
1874+
"learning_rate": {"type": "float", "default": 0.0001},
1875+
})
1876+
fu._validate_hyperparameter_values(options.to_dict(), options) # no raise

sagemaker-train/tests/unit/train/test_common.py

Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -56,6 +56,40 @@ def test_to_dict_all_none_returns_empty(self):
5656
assert result == {}
5757

5858

59+
class TestFineTuningOptionsRequiredKeys:
60+
"""Tests for FineTuningOptions.required_keys()."""
61+
62+
def test_returns_only_required_specs(self):
63+
options = FineTuningOptions({
64+
"learning_rate": {"default": 0.001, "type": "float"},
65+
"required_no_default": {"type": "string", "required": True},
66+
"required_with_default": {"default": 5, "type": "integer", "required": True},
67+
})
68+
assert options.required_keys() == {"required_no_default", "required_with_default"}
69+
70+
def test_returns_empty_when_none_required(self):
71+
options = FineTuningOptions({
72+
"learning_rate": {"default": 0.001, "type": "float"},
73+
"epochs": {"default": 3, "type": "integer"},
74+
})
75+
assert options.required_keys() == set()
76+
77+
def test_required_false_is_excluded(self):
78+
options = FineTuningOptions({
79+
"opt_in": {"default": None, "type": "string", "required": False},
80+
})
81+
assert options.required_keys() == set()
82+
83+
def test_required_key_dropped_by_to_dict_is_still_reported(self):
84+
"""A required spec with no default (None) is dropped by to_dict() but
85+
must still be reported by required_keys() so callers can surface it."""
86+
options = FineTuningOptions({
87+
"required_no_default": {"type": "string", "required": True},
88+
})
89+
assert "required_no_default" not in options.to_dict()
90+
assert options.required_keys() == {"required_no_default"}
91+
92+
5993
import pytest
6094

6195

0 commit comments

Comments
 (0)