Skip to content
Open
33 changes: 33 additions & 0 deletions sagemaker-train/src/sagemaker/train/base_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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."""
Expand Down Expand Up @@ -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).
Expand Down
33 changes: 33 additions & 0 deletions sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
30 changes: 30 additions & 0 deletions sagemaker-train/tests/integ/train/test_sft_trainer_integration.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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)
Original file line number Diff line number Diff line change
Expand Up @@ -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)