Skip to content
Merged
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
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
41 changes: 38 additions & 3 deletions sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1451,10 +1451,45 @@ 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


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:
# 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}."
)
return None

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)
Original file line number Diff line number Diff line change
Expand Up @@ -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
41 changes: 41 additions & 0 deletions sagemaker-train/tests/unit/train/test_base_trainer_serverful.py
Original file line number Diff line number Diff line change
Expand Up @@ -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