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
14 changes: 8 additions & 6 deletions src/sagemaker/tuner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down
40 changes: 40 additions & 0 deletions tests/unit/test_tuner.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
import copy
import os
import re
import typing

import pytest
from mock import Mock, patch
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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
Loading