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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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
Original file line numberDiff line numberDiff 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
Original file line numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all
 blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks");
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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
Original file line numberDiff line numberDiff 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
Original file line numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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
Original file line numberDiff line numberDiff 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
Original file line numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length > 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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
Original file line numberDiff line numberDiff 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
Original file line numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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
Original file line numberDiff line numberDiff 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
Original file line numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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
Original file line numberDiff line numberDiff 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
Original file line numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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
Original file line numberDiff line numberDiff 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
Original file line numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
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 numberDiff line numberDiff 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 DownExpand 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 DownExpand 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
Original file line numberDiff line numberDiff 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
Original file line numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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