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
2 changes: 1 addition & 1 deletion sagemaker-train/src/sagemaker/train/model_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
38 changes: 38 additions & 0 deletions sagemaker-train/tests/unit/train/test_model_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"}}

Expand Down
Loading