From 480d3dc210ea109da6598e8e3e7931de680535a1 Mon Sep 17 00:00:00 2001 From: Mergepath Date: Sat, 1 Aug 2026 18:27:17 +0000 Subject: [PATCH] fix: address #5770 - `ModelTrainer` bug with method to load hyperparameters from file for Amazon Nova Recipe Closes #5770 --- .../src/sagemaker/train/model_trainer.py | 2 +- .../tests/unit/train/test_model_trainer.py | 38 +++++++++++++++++++ 2 files changed, 39 insertions(+), 1 deletion(-) diff --git a/sagemaker-train/src/sagemaker/train/model_trainer.py b/sagemaker-train/src/sagemaker/train/model_trainer.py index 48a80af03b..d22350b4f5 100644 --- a/sagemaker-train/src/sagemaker/train/model_trainer.py +++ b/sagemaker-train/src/sagemaker/train/model_trainer.py @@ -1361,7 +1361,7 @@ def from_recipe( ) if is_nova: if hyperparameters and isinstance(hyperparameters, str): - hyperparameters = cls._validate_and_load_hyperparameters_file(hyperparameters) + hyperparameters = cls._validate_and_fetch_hyperparameters_file(hyperparameters) model_trainer_args["hyperparameters"].update(hyperparameters) elif hyperparameters and isinstance(hyperparameters, dict): model_trainer_args["hyperparameters"].update(hyperparameters) diff --git a/sagemaker-train/tests/unit/train/test_model_trainer.py b/sagemaker-train/tests/unit/train/test_model_trainer.py index ce5d208bbc..99e0ed5060 100644 --- a/sagemaker-train/tests/unit/train/test_model_trainer.py +++ b/sagemaker-train/tests/unit/train/test_model_trainer.py @@ -1570,6 +1570,44 @@ def mock_upload_data(path, bucket, key_prefix): ] +def test_nova_recipe_with_hyperparameters_file(modules_session): + recipe_data = { + "run": { + "name": "dummy-model", + "model_type": "amazon.nova", + "model_name_or_path": "dummy-model", + } + } + hyperparameters = { + "custom_parameter": "custom-value", + "custom_int": 5, + } + + with NamedTemporaryFile(suffix=".yaml", delete=False) as recipe, NamedTemporaryFile( + suffix=".json", delete=False + ) as hyperparameters_file: + with open(recipe.name, "w") as file: + yaml.dump(recipe_data, file) + with open(hyperparameters_file.name, "w") as file: + json.dump(hyperparameters, file) + + trainer = ModelTrainer.from_recipe( + training_recipe=recipe.name, + role=DEFAULT_ROLE, + sagemaker_session=modules_session, + compute=DEFAULT_COMPUTE_CONFIG, + training_image=DEFAULT_IMAGE, + hyperparameters=hyperparameters_file.name, + ) + + assert trainer.hyperparameters["base_model"] == "dummy-model" + assert trainer.hyperparameters["custom_parameter"] == "custom-value" + assert trainer.hyperparameters["custom_int"] == 5 + + os.unlink(recipe.name) + os.unlink(hyperparameters_file.name) + + def test_nova_recipe_with_distillation(modules_session): recipe_data = {"training_config": {"distillation_data": "true", "kms_key": "alias/my-kms-key"}}