Uh oh!
There was an error while loading. Please reload this page.
- Notifications
You must be signed in to change notification settings - Fork 1.3k
fix: ModelTrainer and HyperparameterTuner missing environment variables (5613)#5725
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Uh oh!
There was an error while loading. Please reload this page.
Changes from all commits
f1ea9d5dec47ab862ff2d1e8e693File filter
Filter by extension
Conversations
Uh oh!
There was an error while loading. Please reload this page.
Jump to
Uh oh!
There was an error while loading. Please reload this page.
Diff view
Diff view
There are no files selected for viewing
aviruthen marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. |
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -1504,7 +1504,13 @@ def _build_training_job_definition(self, inputs): | ||
| model_trainer.stopping_condition.max_wait_time_in_seconds | ||
| ) | ||
aviruthen marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. aviruthen marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. | ||
| definition = HyperParameterTrainingJobDefinition( | ||
| # Propagate environment variables from ModelTrainer. | ||
| # Only include when it's a dict (even empty); omit otherwise so the | ||
| # Pydantic field stays Unassigned and is excluded during serialization. | ||
| env = model_trainer.environment | ||
aviruthen marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. | ||
| # Build base kwargs for the definition | ||
| definition_kwargs = dict( | ||
| algorithm_specification=algorithm_spec, | ||
| role_arn=model_trainer.role, | ||
| input_data_config=input_data_config if input_data_config else None, | ||
| @@ -1515,10 +1521,11 @@ def _build_training_job_definition(self, inputs): | ||
| enable_managed_spot_training=model_trainer.compute.enable_managed_spot_training, | ||
| ) | ||
| # Pass through environment variables from model_trainer | ||
| env = getattr(model_trainer, "environment", None) | ||
| if env and isinstance(env, dict): | ||
| definition.environment = env | ||
| # Include environment only when it's a dict (including empty). | ||
| if isinstance(env, dict): | ||
| definition_kwargs["environment"] = env | ||
| definition = HyperParameterTrainingJobDefinition(**definition_kwargs) | ||
| # Pass through VPC config from model_trainer | ||
| networking = getattr(model_trainer, "networking", None) | ||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -596,3 +596,73 @@ def test_build_training_job_definition_includes_spot_params(self): | ||
| assert isinstance( | ||
| definition.stopping_condition.max_wait_time_in_seconds, int | ||
aviruthen marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. | ||
| ), "Max wait time should be set" | ||
aviruthen marked this conversation as resolved.
Uh oh!There was an error while loading. Please reload this page. | ||
| def test_build_training_job_definition_includes_environment_variables(self): | ||
| """Test that _build_training_job_definition includes environment variables. | ||
| This test verifies the fix for GitHub issue #5613 where tuning jobs were | ||
| missing environment variables that were set on the ModelTrainer. | ||
| """ | ||
| mock_trainer = _create_mock_model_trainer() | ||
| mock_trainer.environment = { | ||
| "FOO": "bar", | ||
| "RANDOM_STATE": "42", | ||
| } | ||
| tuner = HyperparameterTuner( | ||
| model_trainer=mock_trainer, | ||
| objective_metric_name="accuracy", | ||
| hyperparameter_ranges=_create_single_hp_range(), | ||
| ) | ||
| definition = tuner._build_training_job_definition(None) | ||
| assert definition.environment is not None, "Environment should not be None" | ||
| assert definition.environment == { | ||
| "FOO": "bar", | ||
| "RANDOM_STATE": "42", | ||
| }, "Environment variables should match those set on ModelTrainer" | ||
| def test_build_training_job_definition_with_none_environment(self): | ||
| """Test that _build_training_job_definition handles None environment gracefully. | ||
| When environment is None, it should not be passed to the Pydantic constructor, | ||
| so the field stays as Unassigned (excluded from serialization). | ||
| """ | ||
| from sagemaker.core.utils.utils import Unassigned | ||
| mock_trainer = _create_mock_model_trainer() | ||
| mock_trainer.environment = None | ||
| tuner = HyperparameterTuner( | ||
| model_trainer=mock_trainer, | ||
| objective_metric_name="accuracy", | ||
| hyperparameter_ranges=_create_single_hp_range(), | ||
| ) | ||
| definition = tuner._build_training_job_definition(None) | ||
| assert isinstance(definition.environment, Unassigned), ( | ||
| "Environment should be Unassigned when model_trainer.environment is None" | ||
| ) | ||
| def test_build_training_job_definition_with_empty_environment(self): | ||
| """Test that _build_training_job_definition passes through empty environment. | ||
| An empty dict is valid for the SageMaker API, so we pass it through as-is | ||
| rather than silently converting it to None. | ||
| """ | ||
| mock_trainer = _create_mock_model_trainer() | ||
| mock_trainer.environment = {} | ||
| tuner = HyperparameterTuner( | ||
| model_trainer=mock_trainer, | ||
| objective_metric_name="accuracy", | ||
| hyperparameter_ranges=_create_single_hp_range(), | ||
| ) | ||
| definition = tuner._build_training_job_definition(None) | ||
| assert definition.environment == {}, ( | ||
| "Empty dict environment should be passed through as-is" | ||
| ) | ||
Uh oh!
There was an error while loading. Please reload this page.
Uh oh!
There was an error while loading. Please reload this page.