diff --git a/src/sagemaker/tuner.py b/src/sagemaker/tuner.py index 685927b52f..0329d60db8 100644 --- a/src/sagemaker/tuner.py +++ b/src/sagemaker/tuner.py @@ -594,7 +594,7 @@ def __init__( self, estimator: EstimatorBase, 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", @@ -624,10 +624,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 @@ -1906,8 +1907,9 @@ def create( names as in estimator_dict, and there must be one entry for each estimator in estimator_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 estimator names as in estimator_dict, and there must be one entry for each estimator in estimator_dict. Each value is diff --git a/tests/unit/test_tuner.py b/tests/unit/test_tuner.py index 9b9e34f95f..6a6dbcbbd0 100644 --- a/tests/unit/test_tuner.py +++ b/tests/unit/test_tuner.py @@ -15,6 +15,7 @@ import copy import os import re +import typing import pytest from mock import Mock, patch @@ -42,6 +43,7 @@ create_transfer_learning_tuner, HyperparameterTuner, ) +from sagemaker.workflow.entities import PipelineVariable from sagemaker.workflow.functions import JsonGet, Join from sagemaker.workflow.parameters import ParameterString, ParameterInteger @@ -2230,3 +2232,41 @@ def test_create_tuner_with_grid_search_strategy(): assert tuner is not None assert tuner.max_jobs is None + + +def test_hyperparameter_ranges_annotation_allows_pipeline_variable_keys(): + """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]``. + """ + 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(estimator): + """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( + estimator=estimator, + objective_metric_name=OBJECTIVE_METRIC_NAME, + hyperparameter_ranges={hparam_name: CategoricalParameter([1, 2])}, + ) + + assert hparam_name in tuner._hyperparameter_ranges