From dbba4a1cdac8eb33b6178380b02116cafeba0e50 Mon Sep 17 00:00:00 2001 From: Syed Jafri Date: Wed, 29 Jul 2026 11:34:46 -0700 Subject: [PATCH 1/4] fix: serverful instance type validations + integ tests --- .../src/sagemaker/train/base_trainer.py | 35 +++++++++- .../train/common_utils/finetune_utils.py | 35 +++++++++- .../train/test_sft_trainer_integration.py | 30 +++++++++ .../train/test_sft_trainer_serverful_smtj.py | 65 +++++++++++++++++++ 4 files changed, 163 insertions(+), 2 deletions(-) diff --git a/sagemaker-train/src/sagemaker/train/base_trainer.py b/sagemaker-train/src/sagemaker/train/base_trainer.py index 468dacd610..2ad0f7e179 100644 --- a/sagemaker-train/src/sagemaker/train/base_trainer.py +++ b/sagemaker-train/src/sagemaker/train/base_trainer.py @@ -31,6 +31,7 @@ get_recipe_s3_uri, _validate_hyperparameter_values, _get_smhp_replicas_enum, + _get_smhp_instance_type_enum, ) from sagemaker.train.common_utils.data_utils import validate_data_path_exists from sagemaker.train.common_utils.metrics_visualizer import plot_training_metrics @@ -779,6 +780,22 @@ def _validate_instance_count(self, instance_count, sagemaker_session): ) return smhp_replicas_enum + def _validate_instance_type(self, instance_type, sagemaker_session): + """Validate instance type against allowed values from SMHP recipe.""" + smhp_instance_type_enum = _get_smhp_instance_type_enum( + model_name=self._model_name, + customization_technique=self._customization_technique, + training_type=self.training_type, + sagemaker_session=sagemaker_session, + ) + + if smhp_instance_type_enum and instance_type not in smhp_instance_type_enum: + raise ValueError( + f"Instance type '{instance_type}' is not supported. " + f"Allowed values: {sorted(smhp_instance_type_enum)}." + ) + return smhp_instance_type_enum + @abstractmethod def train(self, input_data_config: List[InputData], wait: bool = True, logs: bool = True, wait_timeout: Optional[int] = None, dry_run: bool = False): """Common training method that calls the specific implementation.""" @@ -892,14 +909,30 @@ def _channel_mount_path(dataset_uri, channel_name): sagemaker_session=sagemaker_session, ) + # Validates instance type using SMHP override spec as SMTJ override spec doesn't contain instance type + smhp_instance_type_enum = self._validate_instance_type(compute.instance_type, sagemaker_session) + if not smhp_instance_type_enum: + logger.warning( + f"SMHP recipe for {self._model_name}/{self._customization_technique} did not provide a " + f"valid instance_type enum. " + "Instance type validation will be skipped." + ) + # Validate instance count against allowed values from SMHP recipe. smhp_replicas_enum = self._validate_instance_count(compute.instance_count, sagemaker_session) + if smhp_replicas_enum: override_spec.setdefault("replicas", {})["enum"] = smhp_replicas_enum if hasattr(self, 'hyperparameters') and hasattr(self.hyperparameters, '_specs'): self.hyperparameters._specs.setdefault("replicas", {})["enum"] = smhp_replicas_enum if not hasattr(self.hyperparameters, 'replicas'): object.__setattr__(self.hyperparameters, 'replicas', compute.instance_count) + else: + logger.warning( + f"SMHP recipe for {self._model_name}/{self._customization_technique} did not provide a " + f"valid replicas enum. " + "Instance count validation will be skipped." + ) # Inject the resolved dataset channel paths so the rendered recipe's # train_files / val_files are non-empty (the container aborts otherwise). @@ -1487,4 +1520,4 @@ def _train_hyperpod(self, training_dataset=None, validation_dataset=None, ) self._latest_training_job = training_job - return job_name + return job_name \ No newline at end of file diff --git a/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py b/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py index 419ef2115d..da9c9e8000 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py @@ -1459,6 +1459,39 @@ def _get_smhp_replicas_enum(model_name: str, customization_technique: str, train return None +def _get_smhp_instance_type_enum(model_name: str, customization_technique: str, training_type, + sagemaker_session, hub_name: Optional[str] = None) -> Optional[list]: + """Fetch the instance_type enum from the SMHP override spec for the same model/technique. + + SMTJ hub content does not include an instance_type enum in its override spec, but + the SMHP recipe for the same configuration does. This function retrieves that + enum so it can be applied to SMTJ recipe validation. + + Returns: + List of valid instance types, or None if unavailable. + """ + try: + _, smhp_override_spec = _get_recipe_entry_and_override_spec( + model_name=model_name, + customization_technique=customization_technique, + training_type=training_type, + sagemaker_session=sagemaker_session, + platform="hyperpod", + hub_name=hub_name, + ) + instance_type_meta = smhp_override_spec.get("instance_type", {}) + enum_val = instance_type_meta.get("enum") + if isinstance(enum_val, list) and enum_val: + return enum_val + except Exception as e: + logger.warning( + f"Could not fetch valid instance types from SMHP recipe for " + f"{model_name}/{customization_technique}: {e}. " + "Instance type validation will be skipped." + ) + return None + + def _extract_recipe_from_helm_template(template_content: str, customization_technique: str = None) -> str: """Extract the training config YAML from a HyperPod Helm chart template. @@ -1789,4 +1822,4 @@ def extract_image_from_hyperpod_template(template_content: str) -> Optional[str] image_match = re.search(image_pattern, template_content, re.MULTILINE) if image_match: return image_match.group(1).strip() - return None + return None \ No newline at end of file diff --git a/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py b/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py index 68446991c4..a47c493746 100644 --- a/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py +++ b/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py @@ -18,6 +18,7 @@ import pytest import boto3 from sagemaker.core.helper.session_helper import Session +from sagemaker.core.training.configs import TrainingJobCompute from sagemaker.train.sft_trainer import SFTTrainer from sagemaker.train.common import TrainingType @@ -135,3 +136,32 @@ def test_sft_trainer_nova_workflow(sagemaker_session_us_east_1): assert training_job.training_job_status == "Completed" assert hasattr(training_job, 'output_model_package_arn') assert training_job.output_model_package_arn is not None + + +def test_sft_trainer_lora_invalid_instance_type_raises(sagemaker_session): + """An unsupported instance type must raise before a job is submitted. + + Based on the ``test_sft_trainer_lora_workflow`` notebook (Llama LORA on + serverful compute in us-west-2). SFTTrainer validates ``instance_type`` + against the allowed enum from the model's recipe, so ``train()`` should + raise a ``ValueError`` rather than launching a training job. + """ + unique_id = f"{int(time.time())}-{random.randint(1000, 9999)}" + + sft_trainer = SFTTrainer( + model="meta-textgeneration-llama-3-2-1b-instruct", + training_type=TrainingType.LORA, + model_package_group="arn:aws:sagemaker:us-west-2:729646638167:model-package-group/sdk-test-finetuned-models", + training_dataset="s3://mc-flows-sdk-testing/input_data/sft/sample_data_256_final.jsonl", + s3_output_path="s3://mc-flows-sdk-testing/output/", + compute=TrainingJobCompute( + instance_type="ml.t3.medium", # unsupported for SFT training + instance_count=1, + ), + accept_eula=True, + base_job_name=f"sft-lora-integ-bad-type-{unique_id}", + sagemaker_session=sagemaker_session, + ) + + with pytest.raises(ValueError, match="Instance type 'ml.t3.medium' is not supported"): + sft_trainer.train(wait=False, dry_run=True) \ No newline at end of file diff --git a/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py b/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py index 1ef6937a36..24e699e07a 100644 --- a/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py +++ b/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py @@ -162,3 +162,68 @@ def test_sft_trainer_serverful_smtj(sagemaker_session_us_east_1, training_resour f"{training_job.training_job_status}" ) logger.info(f"Training job completed successfully: {training_job.training_job_name}") + + +@pytest.mark.us_east_1 +def test_sft_trainer_serverful_smtj_invalid_instance_type_raises( + sagemaker_session_us_east_1, training_resources +): + """An unsupported instance type must raise before a job is submitted. + + SMTJ compute validates ``instance_type`` against the allowed enum from the + model's SMHP recipe. ``ml.t3.medium`` is not a valid training instance for + Nova, so ``train()`` should raise a ``ValueError`` rather than launching a + training job. + """ + unique_id = f"{int(time.time())}-{random.randint(1000, 9999)}" + + sft_trainer = SFTTrainer( + model="nova-textgeneration-lite-v2", + training_type=TrainingType.LORA, + training_dataset=training_resources["training_dataset"], + s3_output_path=training_resources["s3_output_path"], + compute=TrainingJobCompute( + instance_type="ml.t3.medium", # unsupported for Nova training + instance_count=1, + ), + sagemaker_session=sagemaker_session_us_east_1, + base_job_name=f"sft-smtj-integ-bad-type-{unique_id}", + ) + + with pytest.raises(ValueError, match="Instance type 'ml.t3.medium' is not supported"): + sft_trainer.train(wait=False, dry_run=True) + + +@pytest.mark.us_east_1 +def test_sft_trainer_serverful_smtj_invalid_instance_count_raises( + sagemaker_session_us_east_1, training_resources +): + """An unsupported instance count must raise before a job is submitted. + + Uses a valid instance type so validation reaches the instance-count check, + then supplies an out-of-range count. SMTJ compute validates + ``instance_count`` against the allowed replicas enum from the model's SMHP + recipe, so ``train()`` should raise a ``ValueError``. + """ + unique_id = f"{int(time.time())}-{random.randint(1000, 9999)}" + + invalid_instance_count = 9 + + sft_trainer = SFTTrainer( + model="nova-textgeneration-lite-v2", + training_type=TrainingType.LORA, + training_dataset=training_resources["training_dataset"], + s3_output_path=training_resources["s3_output_path"], + compute=TrainingJobCompute( + instance_type="ml.p4d.24xlarge", # valid so count check is reached + instance_count=invalid_instance_count, + ), + sagemaker_session=sagemaker_session_us_east_1, + base_job_name=f"sft-smtj-integ-bad-count-{unique_id}", + ) + + with pytest.raises( + ValueError, + match=f"Node/Instance count '{invalid_instance_count}' is not supported", + ): + sft_trainer.train(wait=False, dry_run=True) \ No newline at end of file From 9d4fd155e6f5a9479bbf6bedbfde8151362d7f85 Mon Sep 17 00:00:00 2001 From: Syed Jafri Date: Wed, 29 Jul 2026 11:42:24 -0700 Subject: [PATCH 2/4] cleanup: add missing newlines to end of files --- sagemaker-train/src/sagemaker/train/base_trainer.py | 3 ++- .../src/sagemaker/train/common_utils/finetune_utils.py | 2 +- .../tests/integ/train/test_sft_trainer_integration.py | 3 ++- .../tests/integ/train/test_sft_trainer_serverful_smtj.py | 2 +- 4 files changed, 6 insertions(+), 4 deletions(-) diff --git a/sagemaker-train/src/sagemaker/train/base_trainer.py b/sagemaker-train/src/sagemaker/train/base_trainer.py index 2ad0f7e179..44d45a6f61 100644 --- a/sagemaker-train/src/sagemaker/train/base_trainer.py +++ b/sagemaker-train/src/sagemaker/train/base_trainer.py @@ -1520,4 +1520,5 @@ def _train_hyperpod(self, training_dataset=None, validation_dataset=None, ) self._latest_training_job = training_job - return job_name \ No newline at end of file + return job_name + \ No newline at end of file diff --git a/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py b/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py index da9c9e8000..7202647b86 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py @@ -1822,4 +1822,4 @@ def extract_image_from_hyperpod_template(template_content: str) -> Optional[str] image_match = re.search(image_pattern, template_content, re.MULTILINE) if image_match: return image_match.group(1).strip() - return None \ No newline at end of file + return None diff --git a/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py b/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py index a47c493746..7065ad3d26 100644 --- a/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py +++ b/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py @@ -164,4 +164,5 @@ def test_sft_trainer_lora_invalid_instance_type_raises(sagemaker_session): ) with pytest.raises(ValueError, match="Instance type 'ml.t3.medium' is not supported"): - sft_trainer.train(wait=False, dry_run=True) \ No newline at end of file + sft_trainer.train(wait=False, dry_run=True) + \ No newline at end of file diff --git a/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py b/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py index 24e699e07a..28020feb9d 100644 --- a/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py +++ b/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py @@ -226,4 +226,4 @@ def test_sft_trainer_serverful_smtj_invalid_instance_count_raises( ValueError, match=f"Node/Instance count '{invalid_instance_count}' is not supported", ): - sft_trainer.train(wait=False, dry_run=True) \ No newline at end of file + sft_trainer.train(wait=False, dry_run=True) From 3bd4219168d04781f7bac22ed3891a9e9d00a771 Mon Sep 17 00:00:00 2001 From: Syed Jafri Date: Wed, 29 Jul 2026 13:17:29 -0700 Subject: [PATCH 3/4] style(sagemaker-train): Fix missing trailing newlines --- sagemaker-train/src/sagemaker/train/base_trainer.py | 1 - .../tests/integ/train/test_sft_trainer_integration.py | 1 - 2 files changed, 2 deletions(-) diff --git a/sagemaker-train/src/sagemaker/train/base_trainer.py b/sagemaker-train/src/sagemaker/train/base_trainer.py index 44d45a6f61..c10a6bf707 100644 --- a/sagemaker-train/src/sagemaker/train/base_trainer.py +++ b/sagemaker-train/src/sagemaker/train/base_trainer.py @@ -1521,4 +1521,3 @@ def _train_hyperpod(self, training_dataset=None, validation_dataset=None, self._latest_training_job = training_job return job_name - \ No newline at end of file diff --git a/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py b/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py index 7065ad3d26..78e301b5a3 100644 --- a/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py +++ b/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py @@ -165,4 +165,3 @@ def test_sft_trainer_lora_invalid_instance_type_raises(sagemaker_session): with pytest.raises(ValueError, match="Instance type 'ml.t3.medium' is not supported"): sft_trainer.train(wait=False, dry_run=True) - \ No newline at end of file From 0f286b62dc9ec64eed9e3a5b0983ad23ee165b6b Mon Sep 17 00:00:00 2001 From: Syed Jafri Date: Wed, 29 Jul 2026 14:22:41 -0700 Subject: [PATCH 4/4] fix(train): Log SMHP enum fetch failure at debug, not warning + added unit tests --- .../train/common_utils/finetune_utils.py | 14 ++++--- .../train/common_utils/test_finetune_utils.py | 40 ++++++++++++++++++ .../unit/train/test_base_trainer_serverful.py | 41 +++++++++++++++++++ 3 files changed, 89 insertions(+), 6 deletions(-) diff --git a/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py b/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py index 7202647b86..94f5f7af3e 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py @@ -1451,10 +1451,11 @@ def _get_smhp_replicas_enum(model_name: str, customization_technique: str, train if isinstance(enum_val, list) and enum_val: return enum_val except Exception as e: - logger.warning( + # Caller emits the user-facing warning when None is returned; keep the + # exception detail at debug level to avoid a duplicate warning. + logger.debug( f"Could not fetch valid instance counts from SMHP recipe for " - f"{model_name}/{customization_technique}: {e}. " - "Instance count validation will be skipped." + f"{model_name}/{customization_technique}: {e}." ) return None @@ -1484,10 +1485,11 @@ def _get_smhp_instance_type_enum(model_name: str, customization_technique: str, if isinstance(enum_val, list) and enum_val: return enum_val except Exception as e: - logger.warning( + # Caller emits the user-facing warning when None is returned; keep the + # exception detail at debug level to avoid a duplicate warning. + logger.debug( f"Could not fetch valid instance types from SMHP recipe for " - f"{model_name}/{customization_technique}: {e}. " - "Instance type validation will be skipped." + f"{model_name}/{customization_technique}: {e}." ) return None diff --git a/sagemaker-train/tests/unit/train/common_utils/test_finetune_utils.py b/sagemaker-train/tests/unit/train/common_utils/test_finetune_utils.py index 1786479a4d..beb0b877af 100644 --- a/sagemaker-train/tests/unit/train/common_utils/test_finetune_utils.py +++ b/sagemaker-train/tests/unit/train/common_utils/test_finetune_utils.py @@ -1425,3 +1425,43 @@ def test_uses_shared_regex_from_reward_verifier(self): # Both call sites must share the same compiled pattern, not copies. from sagemaker.train.common_utils import rlvr_reward_verifier assert fu.LAMBDA_ARN_REGEX is rlvr_reward_verifier.LAMBDA_ARN_REGEX + + +class TestGetSmhpInstanceTypeEnum: + """Unit tests for _get_smhp_instance_type_enum (SMHP override-spec enum lookup).""" + + def _call(self): + return fu._get_smhp_instance_type_enum( + model_name="my-model", + customization_technique="SFT", + training_type=TrainingType.LORA, + sagemaker_session=MagicMock(), + ) + + @patch.object(fu, "_get_recipe_entry_and_override_spec") + def test_returns_enum_when_present(self, mock_spec): + mock_spec.return_value = ( + {}, + {"instance_type": {"enum": ["ml.p5.48xlarge", "ml.p4d.24xlarge"]}}, + ) + assert self._call() == ["ml.p5.48xlarge", "ml.p4d.24xlarge"] + + @patch.object(fu, "_get_recipe_entry_and_override_spec") + def test_returns_none_when_enum_missing(self, mock_spec): + mock_spec.return_value = ({}, {"instance_type": {}}) + assert self._call() is None + + @patch.object(fu, "_get_recipe_entry_and_override_spec") + def test_returns_none_when_instance_type_key_absent(self, mock_spec): + mock_spec.return_value = ({}, {}) + assert self._call() is None + + @patch.object(fu, "_get_recipe_entry_and_override_spec") + def test_returns_none_when_enum_empty_list(self, mock_spec): + mock_spec.return_value = ({}, {"instance_type": {"enum": []}}) + assert self._call() is None + + @patch.object(fu, "_get_recipe_entry_and_override_spec") + def test_returns_none_on_exception(self, mock_spec): + mock_spec.side_effect = RuntimeError("hub content unavailable") + assert self._call() is None diff --git a/sagemaker-train/tests/unit/train/test_base_trainer_serverful.py b/sagemaker-train/tests/unit/train/test_base_trainer_serverful.py index e0537a9fd8..ab0585eddd 100644 --- a/sagemaker-train/tests/unit/train/test_base_trainer_serverful.py +++ b/sagemaker-train/tests/unit/train/test_base_trainer_serverful.py @@ -318,3 +318,44 @@ def test_mlflow_partial_names_defaulted(self): assert override_spec["mlflow_experiment_name"]["default"] == "user-experiment" assert override_spec["mlflow_run_name"]["default"] == "nova-lite-sft" + + +class TestValidateInstanceType: + """Unit tests for BaseTrainer._validate_instance_type. + + Validates instance types against the SMHP override-spec enum, and skips + validation (returning None) when the enum is unavailable. + """ + + def _trainer(self): + trainer = _ConcreteTrainer.__new__(_ConcreteTrainer) + trainer._model_name = "nova-lite" + trainer._customization_technique = "sft" + trainer.training_type = "lora" + return trainer + + @patch("sagemaker.train.base_trainer._get_smhp_instance_type_enum") + def test_allowed_instance_type_returns_enum(self, mock_enum): + mock_enum.return_value = ["ml.p4d.24xlarge", "ml.p5.48xlarge"] + trainer = self._trainer() + + result = trainer._validate_instance_type("ml.p4d.24xlarge", MagicMock()) + + assert result == ["ml.p4d.24xlarge", "ml.p5.48xlarge"] + + @patch("sagemaker.train.base_trainer._get_smhp_instance_type_enum") + def test_disallowed_instance_type_raises(self, mock_enum): + mock_enum.return_value = ["ml.p4d.24xlarge", "ml.p5.48xlarge"] + trainer = self._trainer() + + with pytest.raises(ValueError, match="is not supported"): + trainer._validate_instance_type("ml.g5.xlarge", MagicMock()) + + @patch("sagemaker.train.base_trainer._get_smhp_instance_type_enum") + def test_skips_validation_when_enum_unavailable(self, mock_enum): + """When the enum can't be fetched, validation is skipped (returns None).""" + mock_enum.return_value = None + trainer = self._trainer() + + # Any instance type is accepted; the method returns None to signal skip. + assert trainer._validate_instance_type("ml.g5.xlarge", MagicMock()) is None