From 003b5499c146d2b8a89c34dd3526c3e62ad03460 Mon Sep 17 00:00:00 2001 From: Mohamed Zeidan Date: Thu, 24 Sep 2026 13:36:39 -0700 Subject: [PATCH] fix: allow PipelineVariable keys in HyperparameterTuner hyperparameter_ranges annotation hyperparameter_ranges was annotated Dict[str, ParameterRange], but the tuner accepts a PipelineVariable (e.g. a pipeline ParameterString) as a dict key (hyperparameter name). mypy therefore reported a false dict-item error for valid code. Broaden the key type to Union[str, PipelineVariable] and align the docstrings. Runtime behavior is unchanged (annotations are not enforced). Fixes #5243 --- sagemaker-train/src/sagemaker/train/tuner.py | 14 +++--- .../tests/unit/train/test_tuner.py | 48 +++++++++++++++++++ 2 files changed, 56 insertions(+), 6 deletions(-) diff --git a/sagemaker-train/src/sagemaker/train/tuner.py b/sagemaker-train/src/sagemaker/train/tuner.py index ed872dc894..ddbd22e315 100644 --- a/sagemaker-train/src/sagemaker/train/tuner.py +++ b/sagemaker-train/src/sagemaker/train/tuner.py @@ -92,7 +92,7 @@ def __init__( self, model_trainer: "ModelTrainer", objective_metric_name: Union[str, PipelineVariable], - hyperparameter_ranges: Dict[str, ParameterRange], + hyperparameter_ranges: Dict[Union[str, PipelineVariable], ParameterRange], metric_definitions: Optional[List[Dict[str, Union[str, PipelineVariable]]]] = None, strategy: Union[str, PipelineVariable] = "Bayesian", objective_type: Union[str, PipelineVariable] = "Maximize", @@ -122,10 +122,11 @@ def __init__( instance. objective_metric_name (str or PipelineVariable): Name of the metric for evaluating training jobs. - hyperparameter_ranges (dict[str, sagemaker.parameter.ParameterRange]): Dictionary of - parameter ranges. These parameter ranges can be one + hyperparameter_ranges (dict[str or PipelineVariable, sagemaker.parameter.ParameterRange]): + Dictionary of parameter ranges. These parameter ranges can be one of three types: Continuous, Integer, or Categorical. The keys of - the dictionary are the names of the hyperparameter, and the + the dictionary are the names of the hyperparameter (a str, or a + PipelineVariable such as a pipeline ParameterString), and the values are the appropriate parameter range class to represent the range. metric_definitions (list[dict[str, str] or list[dict[str, PipelineVariable]]): A list of @@ -1039,8 +1040,9 @@ def create( names as in model_trainer_dict, and there must be one entry for each model_trainer in model_trainer_dict. Each value is a dictionary of sagemaker.parameter.ParameterRange instance, which can be one of three types: Continuous, Integer, or Categorical. - The keys of each ParameterRange dictionaries are the names of the hyperparameter, - and the values are the appropriate parameter range class to represent the range. + The keys of each ParameterRange dictionary are the names of the hyperparameter + (a str, or a PipelineVariable such as a pipeline ParameterString), and the values + are the appropriate parameter range class to represent the range. metric_definitions_dict (dict(str, list[dict]]): Dictionary of metric definitions. The keys are the same set or a subset of model_trainer names as in model_trainer_dict, and there must be one entry for each model_trainer in model_trainer_dict. Each value is diff --git a/sagemaker-train/tests/unit/train/test_tuner.py b/sagemaker-train/tests/unit/train/test_tuner.py index d8010fa2d0..de72ae8692 100644 --- a/sagemaker-train/tests/unit/train/test_tuner.py +++ b/sagemaker-train/tests/unit/train/test_tuner.py @@ -14,6 +14,8 @@ from __future__ import absolute_import +import typing + import pytest from unittest.mock import MagicMock, patch @@ -26,7 +28,10 @@ CategoricalParameter, ContinuousParameter, IntegerParameter, + ParameterRange, ) +from sagemaker.core.helper.pipeline_variable import PipelineVariable +from sagemaker.core.workflow.parameters import ParameterString from sagemaker.core.shapes import ( HyperParameterTuningJobWarmStartConfig, Channel, @@ -148,6 +153,49 @@ def test_init_with_basic_params(self, mock_model_trainer, hyperparameter_ranges) assert tuner.max_jobs == 1 assert tuner.max_parallel_jobs == 1 + def test_hyperparameter_ranges_annotation_allows_pipeline_variable_keys(self): + """Regression for #5243. + + The ``hyperparameter_ranges`` key type must allow PipelineVariable (e.g. a pipeline + ParameterString), not only ``str``, since the tuner accepts pipeline variables as + hyperparameter names. This asserts the resolved annotation, so it fails if the type + is narrowed back to ``Dict[str, ParameterRange]``. + """ + # Read the raw annotation object directly: get_type_hints would fail resolving the + # TYPE_CHECKING-only "ModelTrainer" forward ref, and this module does not use + # ``from __future__ import annotations``, so this is a real typing object. + annotation = HyperparameterTuner.__init__.__annotations__["hyperparameter_ranges"] + + key_type, value_type = typing.get_args(annotation) + key_options = typing.get_args(key_type) # (str, PipelineVariable) + + assert str in key_options, f"str must remain a valid key type, got {key_options}" + assert ( + PipelineVariable in key_options + ), f"PipelineVariable must be an allowed key type, got {key_options}" + assert value_type is ParameterRange + + def test_init_with_pipeline_variable_hyperparameter_key(self, mock_model_trainer): + """A pipeline ParameterString used as a hyperparameter-range key is accepted (#5243). + + NOTE: this is a runtime sanity check, not the regression guard -- annotations are not + enforced at runtime, so this passes with or without the fix. The guard against + re-narrowing the type is test_hyperparameter_ranges_annotation_allows_pipeline_variable_keys. + """ + hparam_name = ParameterString(name="HParamName", default_value="hparam") + + tuner = HyperparameterTuner( + model_trainer=mock_model_trainer, + objective_metric_name="valid:loss", + objective_type="Minimize", + hyperparameter_ranges={hparam_name: CategoricalParameter([1, 2])}, + strategy=GRID_SEARCH, + max_jobs=2, + max_parallel_jobs=1, + ) + + assert hparam_name in tuner._hyperparameter_ranges + def test_init_with_custom_strategy(self, mock_model_trainer, hyperparameter_ranges): """Test initialization with custom strategy.""" tuner = HyperparameterTuner(