From 91bdacd7e3fe5f4939ad3d37985048f0cbba8b88 Mon Sep 17 00:00:00 2001 From: Syed Jafri Date: Wed, 22 Jul 2026 17:31:51 -0700 Subject: [PATCH 1/7] fix: stream_logs_smhp extract training job from obj --- sagemaker-train/src/sagemaker/train/base_trainer.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/sagemaker-train/src/sagemaker/train/base_trainer.py b/sagemaker-train/src/sagemaker/train/base_trainer.py index d3c5957c12..468dacd610 100644 --- a/sagemaker-train/src/sagemaker/train/base_trainer.py +++ b/sagemaker-train/src/sagemaker/train/base_trainer.py @@ -701,7 +701,12 @@ def _stream_logs_smtj(self, training_job, poll: int) -> None: def _stream_logs_smhp(self, training_job, compute, poll: int, start_time_ms=None) -> None: """Stream logs for a HyperPod job using filter_log_events polling.""" - job_id = training_job if isinstance(training_job, str) else str(training_job) + if isinstance(training_job, str): + job_id = training_job + elif hasattr(training_job, 'training_job_name'): + job_id = training_job.training_job_name + else: + job_id = str(training_job) sagemaker_session = TrainDefaults.get_sagemaker_session( sagemaker_session=self.sagemaker_session From 5b45b3c85a32e3501b1cfc01c118484ac1aef734 Mon Sep 17 00:00:00 2001 From: Syed Jafri Date: Fri, 24 Jul 2026 15:41:42 -0700 Subject: [PATCH 2/7] fix: render mlflow metrics as png to handle large number of metrics --- .../train/common_utils/metrics_visualizer.py | 52 ++++++++++++++----- 1 file changed, 40 insertions(+), 12 deletions(-) diff --git a/sagemaker-train/src/sagemaker/train/common_utils/metrics_visualizer.py b/sagemaker-train/src/sagemaker/train/common_utils/metrics_visualizer.py index fe837a91fc..cf12afb442 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/metrics_visualizer.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/metrics_visualizer.py @@ -231,13 +231,24 @@ def plot_training_metrics( history = client.get_metric_history(run_id, metric_name) if history: metric_data[metric_name] = history - - # Plot + num_metrics = len(metric_data) + if num_metrics == 0: + logger.warning("No metric history found for run %s. Nothing to plot.", run_id) + return + + # Build the figure at full size (2 columns, all metrics). The inline + # backend can't rasterize this when it gets too tall, so instead of + # calling display(fig) directly we render to an in-memory PNG and embed + # it inside a scrollable HTML
. This lets the notebook display all + # metrics without choking on pixel-count limits. + import io + import base64 + rows = (num_metrics + 1) // 2 - fig, axes = plt.subplots(rows, 2, figsize=(figsize[0], figsize[1] * rows)) - axes = axes.flatten() if num_metrics > 1 else [axes] - + fig, axes = plt.subplots(rows, 2, figsize=(figsize[0], figsize[1] * rows), squeeze=False) + axes = axes.flatten() + for idx, (metric_name, history) in enumerate(metric_data.items()): steps = [h.step for h in history] values = [h.value for h in history] @@ -246,14 +257,31 @@ def plot_training_metrics( axes[idx].set_ylabel('Value') axes[idx].set_title(metric_name, fontweight='bold') axes[idx].grid(True, alpha=0.3) - - for idx in range(len(metric_data), len(axes)): + + for idx in range(num_metrics, len(axes)): axes[idx].set_visible(False) - - plt.suptitle(f'Training Metrics: {training_job.training_job_name}', fontweight='bold', fontsize=14) - plt.tight_layout(rect=[0, 0, 1, 0.98]) # Leave small space for suptitle - display(fig) - plt.close() + + fig.suptitle( + f'Training Metrics: {training_job.training_job_name}', + fontweight='bold', fontsize=14, + ) + fig.tight_layout(rect=[0, 0, 1, 0.98]) + + # Render to PNG buffer for inline display + buf = io.BytesIO() + fig.savefig(buf, format="png", dpi=80, bbox_inches="tight") + plt.close(fig) + buf.seek(0) + + # Embed as a scrollable HTML image in the notebook + b64 = base64.b64encode(buf.getvalue()).decode() + from IPython.display import HTML + display(HTML( + f'
' + f'' + f'
' + )) def get_available_metrics(training_job: TrainingJob) -> List[str]: From 0ec41c02b76b9e234929c93bc5768ac7e253bfa3 Mon Sep 17 00:00:00 2001 From: Syed Jafri Date: Mon, 27 Jul 2026 13:15:22 -0700 Subject: [PATCH 3/7] code cleanup: move io, base64 to top level imports --- .../src/sagemaker/train/common_utils/metrics_visualizer.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/sagemaker-train/src/sagemaker/train/common_utils/metrics_visualizer.py b/sagemaker-train/src/sagemaker/train/common_utils/metrics_visualizer.py index cf12afb442..ec5bb2dc0a 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/metrics_visualizer.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/metrics_visualizer.py @@ -1,7 +1,10 @@ """MLflow metrics visualization utilities for SageMaker training jobs.""" +import io +import base64 import logging from typing import Optional, List, Dict, Any + from sagemaker.core.resources import TrainingJob logger = logging.getLogger(__name__) @@ -209,7 +212,6 @@ def plot_training_metrics( import mlflow from mlflow.tracking import MlflowClient from IPython.display import display - import logging logging.getLogger('botocore.credentials').setLevel(logging.WARNING) @@ -242,9 +244,6 @@ def plot_training_metrics( # calling display(fig) directly we render to an in-memory PNG and embed # it inside a scrollable HTML
. This lets the notebook display all # metrics without choking on pixel-count limits. - import io - import base64 - rows = (num_metrics + 1) // 2 fig, axes = plt.subplots(rows, 2, figsize=(figsize[0], figsize[1] * rows), squeeze=False) axes = axes.flatten() From 3b4be795d71f3a7b6778787935f47d0d6923a0ed Mon Sep 17 00:00:00 2001 From: Syed Jafri Date: Mon, 27 Jul 2026 16:03:20 -0700 Subject: [PATCH 4/7] fix: add helper method to create sns topic --- .../train/common_utils/notifications.py | 95 +++++++++++++++++-- 1 file changed, 87 insertions(+), 8 deletions(-) diff --git a/sagemaker-train/src/sagemaker/train/common_utils/notifications.py b/sagemaker-train/src/sagemaker/train/common_utils/notifications.py index c7e65c2792..c9990fcb73 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/notifications.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/notifications.py @@ -150,6 +150,85 @@ def _validate_sns_topic(sns_client, topic_arn: str) -> None: raise +def _build_topic_policy(topic_arn: str, account_id: str) -> str: + """Build the SNS topic access policy for EventBridge notifications. + + Includes the default owner statement (standard SNS actions scoped to the + topic's own account via ``AWS:SourceAccount``) plus an explicit statement + granting the EventBridge service principal ``SNS:Publish``. + + Args: + topic_arn: The ARN of the SNS topic. + account_id: The AWS account ID that owns the topic. + + Returns: + JSON string of the topic access policy. + """ + policy = { + "Version": "2008-10-17", + "Id": "__default_policy_ID", + "Statement": [ + { + "Sid": "__default_statement_ID", + "Effect": "Allow", + "Principal": {"AWS": "*"}, + "Action": [ + "SNS:Publish", + "SNS:RemovePermission", + "SNS:SetTopicAttributes", + "SNS:DeleteTopic", + "SNS:ListSubscriptionsByTopic", + "SNS:GetTopicAttributes", + "SNS:AddPermission", + "SNS:Subscribe", + ], + "Resource": topic_arn, + "Condition": {"StringEquals": {"AWS:SourceAccount": account_id}}, + }, + { + "Sid": "AllowEventBridgePublish", + "Effect": "Allow", + "Principal": {"Service": "events.amazonaws.com"}, + "Action": "SNS:Publish", + "Resource": topic_arn, + "Condition": {"StringEquals": {"AWS:SourceAccount": account_id}}, + }, + ], + } + return json.dumps(policy) + + +def create_notification_topic(topic_name: str, sagemaker_session) -> str: + """Create an SNS topic configured for EventBridge notifications. + + Creates a new SNS topic and sets an access policy that allows EventBridge + in the same account to publish to it. The returned ARN can be passed + straight to :func:`enable_notifications`. + + Args: + topic_name: Name of the SNS topic to create. + sagemaker_session: SageMaker session (provides boto_session). + + Returns: + The ARN of the created SNS topic. + """ + region_name = sagemaker_session.boto_session.region_name + sns_client = sagemaker_session.boto_session.client("sns", region_name=region_name) + + response = sns_client.create_topic(Name=topic_name) + topic_arn = response["TopicArn"] + # topic ARN: arn:aws:sns::: + account_id = topic_arn.split(":")[4] + + policy = _build_topic_policy(topic_arn, account_id) + sns_client.set_topic_attributes( + TopicArn=topic_arn, AttributeName="Policy", AttributeValue=policy + ) + + logger.info(f"Created SNS topic with EventBridge publish policy: {topic_arn}") + return topic_arn + + def enable_notifications( sns_topic_arn: str, sagemaker_session, @@ -240,14 +319,14 @@ def enable_notifications( try: events_client.put_targets(**put_targets_kwargs) - except Exception as e: - raise ValueError( - f"Failed to attach SNS topic to EventBridge rule. " - f"Ensure your SNS topic has a resource policy allowing " - f"EventBridge to publish to it. Add this to the topic's access policy:\n" - f'{{"Effect": "Allow", "Principal": {{"Service": "events.amazonaws.com"}}, ' - f'"Action": "sns:Publish", "Resource": "{sns_topic_arn}"}}' - ) from e + except ClientError as e: + error_code = e.response.get("Error", {}).get("Code", "") + if error_code in ("AccessDeniedException", "AccessDenied"): + raise PermissionError( + "Missing permission events:PutTargets to attach the SNS topic " + "to the EventBridge rule." + ) from e + raise logger.info(f"Notifications enabled: {normalized_events} -> {sns_topic_arn}") return rule_arn From ce4cebed73ad4669b8b2fb125b273b8d0ca6d26a Mon Sep 17 00:00:00 2001 From: Syed Jafri Date: Tue, 28 Jul 2026 12:28:02 -0700 Subject: [PATCH 5/7] fix: renamed IDs for readability in SNS access policy --- .../src/sagemaker/train/common_utils/notifications.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/sagemaker-train/src/sagemaker/train/common_utils/notifications.py b/sagemaker-train/src/sagemaker/train/common_utils/notifications.py index c9990fcb73..eced509685 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/notifications.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/notifications.py @@ -166,10 +166,10 @@ def _build_topic_policy(topic_arn: str, account_id: str) -> str: """ policy = { "Version": "2008-10-17", - "Id": "__default_policy_ID", + "Id": "SageMakerNotificationsTopicPolicy", "Statement": [ { - "Sid": "__default_statement_ID", + "Sid": "SNSTopicAdministration", "Effect": "Allow", "Principal": {"AWS": "*"}, "Action": [ From 8d2c1619d619f95fd38712757a70e75e871e5918 Mon Sep 17 00:00:00 2001 From: Syed Jafri Date: Tue, 28 Jul 2026 16:33:52 -0700 Subject: [PATCH 6/7] fix: instance type/count validations for smtj serverful --- .../src/sagemaker/train/base_trainer.py | 33 ++++++++++ .../train/common_utils/finetune_utils.py | 33 ++++++++++ .../train/test_sft_trainer_integration.py | 30 +++++++++ .../train/test_sft_trainer_serverful_smtj.py | 65 +++++++++++++++++++ 4 files changed, 161 insertions(+) diff --git a/sagemaker-train/src/sagemaker/train/base_trainer.py b/sagemaker-train/src/sagemaker/train/base_trainer.py index 468dacd610..c10a6bf707 100644 --- a/sagemaker-train/src/sagemaker/train/base_trainer.py +++ b/sagemaker-train/src/sagemaker/train/base_trainer.py @@ -31,6 +31,7 @@ get_recipe_s3_uri, _validate_hyperparameter_values, _get_smhp_replicas_enum, + _get_smhp_instance_type_enum, ) from sagemaker.train.common_utils.data_utils import validate_data_path_exists from sagemaker.train.common_utils.metrics_visualizer import plot_training_metrics @@ -779,6 +780,22 @@ def _validate_instance_count(self, instance_count, sagemaker_session): ) return smhp_replicas_enum + def _validate_instance_type(self, instance_type, sagemaker_session): + """Validate instance type against allowed values from SMHP recipe.""" + smhp_instance_type_enum = _get_smhp_instance_type_enum( + model_name=self._model_name, + customization_technique=self._customization_technique, + training_type=self.training_type, + sagemaker_session=sagemaker_session, + ) + + if smhp_instance_type_enum and instance_type not in smhp_instance_type_enum: + raise ValueError( + f"Instance type '{instance_type}' is not supported. " + f"Allowed values: {sorted(smhp_instance_type_enum)}." + ) + return smhp_instance_type_enum + @abstractmethod def train(self, input_data_config: List[InputData], wait: bool = True, logs: bool = True, wait_timeout: Optional[int] = None, dry_run: bool = False): """Common training method that calls the specific implementation.""" @@ -892,14 +909,30 @@ def _channel_mount_path(dataset_uri, channel_name): sagemaker_session=sagemaker_session, ) + # Validates instance type using SMHP override spec as SMTJ override spec doesn't contain instance type + smhp_instance_type_enum = self._validate_instance_type(compute.instance_type, sagemaker_session) + if not smhp_instance_type_enum: + logger.warning( + f"SMHP recipe for {self._model_name}/{self._customization_technique} did not provide a " + f"valid instance_type enum. " + "Instance type validation will be skipped." + ) + # Validate instance count against allowed values from SMHP recipe. smhp_replicas_enum = self._validate_instance_count(compute.instance_count, sagemaker_session) + if smhp_replicas_enum: override_spec.setdefault("replicas", {})["enum"] = smhp_replicas_enum if hasattr(self, 'hyperparameters') and hasattr(self.hyperparameters, '_specs'): self.hyperparameters._specs.setdefault("replicas", {})["enum"] = smhp_replicas_enum if not hasattr(self.hyperparameters, 'replicas'): object.__setattr__(self.hyperparameters, 'replicas', compute.instance_count) + else: + logger.warning( + f"SMHP recipe for {self._model_name}/{self._customization_technique} did not provide a " + f"valid replicas enum. " + "Instance count validation will be skipped." + ) # Inject the resolved dataset channel paths so the rendered recipe's # train_files / val_files are non-empty (the container aborts otherwise). diff --git a/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py b/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py index 756982375f..405c7323ba 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py @@ -1418,6 +1418,39 @@ def _get_smhp_replicas_enum(model_name: str, customization_technique: str, train return None +def _get_smhp_instance_type_enum(model_name: str, customization_technique: str, training_type, + sagemaker_session, hub_name: Optional[str] = None) -> Optional[list]: + """Fetch the instance_type enum from the SMHP override spec for the same model/technique. + + SMTJ hub content does not include an instance_type enum in its override spec, but + the SMHP recipe for the same configuration does. This function retrieves that + enum so it can be applied to SMTJ recipe validation. + + Returns: + List of valid instance types, or None if unavailable. + """ + try: + _, smhp_override_spec = _get_recipe_entry_and_override_spec( + model_name=model_name, + customization_technique=customization_technique, + training_type=training_type, + sagemaker_session=sagemaker_session, + platform="hyperpod", + hub_name=hub_name, + ) + instance_type_meta = smhp_override_spec.get("instance_type", {}) + enum_val = instance_type_meta.get("enum") + if isinstance(enum_val, list) and enum_val: + return enum_val + except Exception as e: + logger.warning( + f"Could not fetch valid instance types from SMHP recipe for " + f"{model_name}/{customization_technique}: {e}. " + "Instance type validation will be skipped." + ) + return None + + def _extract_recipe_from_helm_template(template_content: str, customization_technique: str = None) -> str: """Extract the training config YAML from a HyperPod Helm chart template. diff --git a/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py b/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py index 68446991c4..78e301b5a3 100644 --- a/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py +++ b/sagemaker-train/tests/integ/train/test_sft_trainer_integration.py @@ -18,6 +18,7 @@ import pytest import boto3 from sagemaker.core.helper.session_helper import Session +from sagemaker.core.training.configs import TrainingJobCompute from sagemaker.train.sft_trainer import SFTTrainer from sagemaker.train.common import TrainingType @@ -135,3 +136,32 @@ def test_sft_trainer_nova_workflow(sagemaker_session_us_east_1): assert training_job.training_job_status == "Completed" assert hasattr(training_job, 'output_model_package_arn') assert training_job.output_model_package_arn is not None + + +def test_sft_trainer_lora_invalid_instance_type_raises(sagemaker_session): + """An unsupported instance type must raise before a job is submitted. + + Based on the ``test_sft_trainer_lora_workflow`` notebook (Llama LORA on + serverful compute in us-west-2). SFTTrainer validates ``instance_type`` + against the allowed enum from the model's recipe, so ``train()`` should + raise a ``ValueError`` rather than launching a training job. + """ + unique_id = f"{int(time.time())}-{random.randint(1000, 9999)}" + + sft_trainer = SFTTrainer( + model="meta-textgeneration-llama-3-2-1b-instruct", + training_type=TrainingType.LORA, + model_package_group="arn:aws:sagemaker:us-west-2:729646638167:model-package-group/sdk-test-finetuned-models", + training_dataset="s3://mc-flows-sdk-testing/input_data/sft/sample_data_256_final.jsonl", + s3_output_path="s3://mc-flows-sdk-testing/output/", + compute=TrainingJobCompute( + instance_type="ml.t3.medium", # unsupported for SFT training + instance_count=1, + ), + accept_eula=True, + base_job_name=f"sft-lora-integ-bad-type-{unique_id}", + sagemaker_session=sagemaker_session, + ) + + with pytest.raises(ValueError, match="Instance type 'ml.t3.medium' is not supported"): + sft_trainer.train(wait=False, dry_run=True) diff --git a/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py b/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py index 1ef6937a36..b5574edf4f 100644 --- a/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py +++ b/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py @@ -162,3 +162,68 @@ def test_sft_trainer_serverful_smtj(sagemaker_session_us_east_1, training_resour f"{training_job.training_job_status}" ) logger.info(f"Training job completed successfully: {training_job.training_job_name}") + + +@pytest.mark.us_east_1 +def test_sft_trainer_serverful_smtj_invalid_instance_type_raises( + sagemaker_session_us_east_1, training_resources +): + """An unsupported instance type must raise before a job is submitted. + + SMTJ compute validates ``instance_type`` against the allowed enum from the + model's SMHP recipe. ``ml.t3.medium`` is not a valid training instance for + Nova, so ``train()`` should raise a ``ValueError`` rather than launching a + training job. + """ + unique_id = f"{int(time.time())}-{random.randint(1000, 9999)}" + + sft_trainer = SFTTrainer( + model="nova-textgeneration-lite-v2", + training_type=TrainingType.LORA, + training_dataset=training_resources["training_dataset"], + s3_output_path=training_resources["s3_output_path"], + compute=TrainingJobCompute( + instance_type="ml.t3.medium", # unsupported for Nova training + instance_count=1, + ), + sagemaker_session=sagemaker_session_us_east_1, + base_job_name=f"sft-smtj-integ-bad-type-{unique_id}", + ) + + with pytest.raises(ValueError, match="Instance type 'ml.t3.medium' is not supported"): + sft_trainer.train(wait=False, dry_run=True) + + +@pytest.mark.us_east_1 +def test_sft_trainer_serverful_smtj_invalid_instance_count_raises( + sagemaker_session_us_east_1, training_resources +): + """An unsupported instance count must raise before a job is submitted. + + Uses a valid instance type so validation reaches the instance-count check, + then supplies an out-of-range count. SMTJ compute validates + ``instance_count`` against the allowed replicas enum from the model's SMHP + recipe, so ``train()`` should raise a ``ValueError``. + """ + unique_id = f"{int(time.time())}-{random.randint(1000, 9999)}" + + invalid_instance_count = 9 + + sft_trainer = SFTTrainer( + model="amazon.nova-micro-v1", + training_type=TrainingType.LORA, + training_dataset=training_resources["training_dataset"], + s3_output_path=training_resources["s3_output_path"], + compute=TrainingJobCompute( + instance_type="ml.g5.12xlarge", # 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) From a0613e72a8518af17255b285e2b90e7eddc0846b Mon Sep 17 00:00:00 2001 From: Syed Jafri Date: Tue, 28 Jul 2026 16:42:12 -0700 Subject: [PATCH 7/7] fix: use nova lite v2 for SMTJ serverful instance count validation test --- .../tests/integ/train/test_sft_trainer_serverful_smtj.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py b/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py index b5574edf4f..28020feb9d 100644 --- a/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py +++ b/sagemaker-train/tests/integ/train/test_sft_trainer_serverful_smtj.py @@ -210,12 +210,12 @@ def test_sft_trainer_serverful_smtj_invalid_instance_count_raises( invalid_instance_count = 9 sft_trainer = SFTTrainer( - model="amazon.nova-micro-v1", + 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.g5.12xlarge", # valid so count check is reached + instance_type="ml.p4d.24xlarge", # valid so count check is reached instance_count=invalid_instance_count, ), sagemaker_session=sagemaker_session_us_east_1,