diff --git a/sagemaker-train/src/sagemaker/train/base_trainer.py b/sagemaker-train/src/sagemaker/train/base_trainer.py index 468dacd610..c10a6bf707 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). 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..7202647b86 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. 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..78e301b5a3 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) 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..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 @@ -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)