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

from __future__ import absolute_import

import typing

import pytest
from unittest.mock import MagicMock, patch

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