From bc2e960aba0dc432bf57170b3186cde158b3a7cc Mon Sep 17 00:00:00 2001 From: jzhaoqwa Date: Mon, 27 Jul 2026 15:10:16 -0700 Subject: [PATCH 01/10] Update error message on ModelBuilder when deploying from S3 checkpoint --- .../src/sagemaker/serve/model_builder.py | 25 +++++++++++++++++++ .../src/sagemaker/train/model_trainer.py | 2 ++ 2 files changed, 27 insertions(+) diff --git a/sagemaker-serve/src/sagemaker/serve/model_builder.py b/sagemaker-serve/src/sagemaker/serve/model_builder.py index 81198b6720..263c78e5eb 100644 --- a/sagemaker-serve/src/sagemaker/serve/model_builder.py +++ b/sagemaker-serve/src/sagemaker/serve/model_builder.py @@ -2853,6 +2853,16 @@ def _build_single_modelbuilder( self.serve_settings = self._get_serve_setting() + # Validate BaseTrainer has a completed training job before proceeding + if isinstance(self.model, BaseTrainer): + if not hasattr(self.model, "_latest_training_job") or self.model._latest_training_job is None: + raise ValueError( + "The trainer passed to ModelBuilder does not have a completed training job. " + "Either call trainer.train() first, or manually set " + "trainer._latest_training_job = TrainingJob.get(training_job_name='') " + "to attach a previously completed job." + ) + # Handle model customization (fine-tuned models) if self._is_model_customization(): if mode is not None and mode != Mode.SAGEMAKER_ENDPOINT: @@ -2900,6 +2910,21 @@ def _build_single_modelbuilder( base_model = model_package.inference_specification.containers[0].base_model if base_model is not None: self._fetch_and_cache_recipe_config() + else: + # No model package available (e.g. serverful SMTJ training job). + # Validate required fields that would normally be auto-resolved. + missing_fields = [] + if not self.image_uri: + missing_fields.append("image_uri") + if isinstance(self.model, BaseTrainer) and not self.model.base_model_name: + missing_fields.append("trainer.base_model_name") + if missing_fields: + raise ValueError( + f"When deploying a model from an S3 checkpoint (e.g. a serverful SMTJ " + f"training job), the following must be provided because no model package " + f"is available to auto-resolve them: {', '.join(missing_fields)}. " + f"Set these on the ModelBuilder or trainer before calling build()." + ) # Nova models use a completely different deployment architecture if self._is_nova_model(): diff --git a/sagemaker-train/src/sagemaker/train/model_trainer.py b/sagemaker-train/src/sagemaker/train/model_trainer.py index 0f77cacd04..3554a74ea7 100644 --- a/sagemaker-train/src/sagemaker/train/model_trainer.py +++ b/sagemaker-train/src/sagemaker/train/model_trainer.py @@ -810,6 +810,8 @@ def train( """ training_request = self._create_training_job_args(input_data_config=input_data_config) + logger.info(f"Training Job Name: {training_request['training_job_name']}") + if dry_run: logger.info("Dry-run validation passed. No job submitted.") return None From a453d5ad1eebef1ed966c485e2f963a12eb9508b Mon Sep 17 00:00:00 2001 From: jzhaoqwa Date: Mon, 27 Jul 2026 16:17:26 -0700 Subject: [PATCH 02/10] Update ModelBuilder to automatically find image_uri --- .../src/sagemaker/serve/model_builder.py | 96 +++++++++++++++---- 1 file changed, 79 insertions(+), 17 deletions(-) diff --git a/sagemaker-serve/src/sagemaker/serve/model_builder.py b/sagemaker-serve/src/sagemaker/serve/model_builder.py index 263c78e5eb..8bb908c676 100644 --- a/sagemaker-serve/src/sagemaker/serve/model_builder.py +++ b/sagemaker-serve/src/sagemaker/serve/model_builder.py @@ -167,6 +167,8 @@ MODEL_SOURCE_TAG_KEY, ) from sagemaker.core.training.utils import resolve_nova_checkpoint_uri +from sagemaker.train.common_utils.model_aliases import normalize_model_name + _LOWEST_MMS_VERSION = "1.2" SCRIPT_PARAM_NAME = "sagemaker_program" @@ -1072,6 +1074,71 @@ def _fetch_and_cache_recipe_config(self): f"Please use a model that supports deployment or contact AWS support for assistance." ) + def _resolve_hosting_config_from_base_model_name(self): + """Resolve image_uri, env_vars, and instance_type from Hub using base_model_name. + + Used when no model package is available (e.g. serverful SMTJ training jobs). + The base_model_name is normalized to a Hub content name, then the Hub document + is fetched to extract hosting configuration — replicating what the model-package + path does via _fetch_and_cache_recipe_config(). + """ + + base_model_name = self._base_model_name() + if not base_model_name: + raise ValueError( + "base_model_name is required when deploying a model from an S3 checkpoint " + "(e.g. a serverful SMTJ training job) because no model package is available " + "to auto-resolve the inference container image. " + "Set trainer.base_model_name before calling build()." + ) + + # If user already provided image_uri, skip hub resolution + if self.image_uri: + logger.info(f"Using provided image_uri: {self.image_uri}") + return + + hub_content_name = normalize_model_name(base_model_name) + hub_name = getattr(self, "hub_name", None) or "SageMakerPublicHub" + + try: + hub_content = HubContent.get( + hub_content_type="Model", + hub_name=hub_name, + hub_content_name=hub_content_name, + ) + hub_document = json.loads(hub_content.hub_content_document) + except Exception as e: + raise ValueError( + f"Could not resolve hosting configuration from Hub for model " + f"'{base_model_name}' (hub_content_name='{hub_content_name}'). " + f"Please provide image_uri explicitly to ModelBuilder. Error: {e}" + ) + + # Try to find hosting configs in the RecipeCollection + for recipe in hub_document.get("RecipeCollection", []): + hosting_configs = recipe.get("HostingConfigs", []) + if hosting_configs: + config = self._select_hosting_config_entry(hosting_configs) + self.image_uri = config.get("EcrAddress") + if self.image_uri: + logger.info(f"Resolved image_uri from Hub: {self.image_uri}") + return + + raise ValueError( + f"Could not resolve inference image URI from Hub for model " + f"'{base_model_name}' (hub_content_name='{hub_content_name}'). " + f"No hosting configuration found in the hub document. " + f"Please provide image_uri explicitly to ModelBuilder." + ) + + @staticmethod + def _select_hosting_config_entry(hosting_configs): + """Select the best hosting config entry, preferring 'Default' profile.""" + return next( + (cfg for cfg in hosting_configs if cfg.get("Profile") == "Default"), + hosting_configs[0], + ) + # Nova escrow ECR accounts per region _NOVA_ESCROW_ACCOUNTS = { "us-east-1": "708977205387", @@ -1267,7 +1334,13 @@ def _get_nova_hosting_config(self, instance_type=None): return hub_config model_package = self._fetch_model_package() - hub_content_name = model_package.inference_specification.containers[0].base_model.hub_content_name + if model_package: + hub_content_name = model_package.inference_specification.containers[0].base_model.hub_content_name + else: + # No model package (e.g. SMTJ trainer): resolve from base_model_name + from sagemaker.train.common_utils.model_aliases import normalize_model_name + base_model_name = self._base_model_name() + hub_content_name = normalize_model_name(base_model_name) if base_model_name else None configs = self._NOVA_HOSTING_CONFIGS.get(hub_content_name) if not configs: @@ -2904,27 +2977,16 @@ def _build_single_modelbuilder( # Fetch recipe config first to set image_uri, instance_type, env_vars, # and s3_upload_path. Only possible when a model package is available; - # trainers built from an S3 checkpoint carry no package, so the caller - # must supply image_uri/instance_type/env_vars directly. + # trainers built from an S3 checkpoint carry no package, so we resolve + # hosting config from the Hub using base_model_name. if model_package is not None: base_model = model_package.inference_specification.containers[0].base_model if base_model is not None: self._fetch_and_cache_recipe_config() else: - # No model package available (e.g. serverful SMTJ training job). - # Validate required fields that would normally be auto-resolved. - missing_fields = [] - if not self.image_uri: - missing_fields.append("image_uri") - if isinstance(self.model, BaseTrainer) and not self.model.base_model_name: - missing_fields.append("trainer.base_model_name") - if missing_fields: - raise ValueError( - f"When deploying a model from an S3 checkpoint (e.g. a serverful SMTJ " - f"training job), the following must be provided because no model package " - f"is available to auto-resolve them: {', '.join(missing_fields)}. " - f"Set these on the ModelBuilder or trainer before calling build()." - ) + # No model package (e.g. serverful SMTJ training job). + # Resolve hosting config from Hub using base_model_name. + self._resolve_hosting_config_from_base_model_name() # Nova models use a completely different deployment architecture if self._is_nova_model(): From dfa6781da1039ac17593072eb286ee194659e2d5 Mon Sep 17 00:00:00 2001 From: jzhaoqwa Date: Tue, 28 Jul 2026 12:29:26 -0700 Subject: [PATCH 03/10] Update import and methods --- .../src/sagemaker/serve/model_builder.py | 28 +++++++++---------- 1 file changed, 13 insertions(+), 15 deletions(-) diff --git a/sagemaker-serve/src/sagemaker/serve/model_builder.py b/sagemaker-serve/src/sagemaker/serve/model_builder.py index 8bb908c676..b36ab609ab 100644 --- a/sagemaker-serve/src/sagemaker/serve/model_builder.py +++ b/sagemaker-serve/src/sagemaker/serve/model_builder.py @@ -1075,28 +1075,18 @@ def _fetch_and_cache_recipe_config(self): ) def _resolve_hosting_config_from_base_model_name(self): - """Resolve image_uri, env_vars, and instance_type from Hub using base_model_name. + """Resolve image_uri from Hub using base_model_name. Used when no model package is available (e.g. serverful SMTJ training jobs). The base_model_name is normalized to a Hub content name, then the Hub document - is fetched to extract hosting configuration — replicating what the model-package - path does via _fetch_and_cache_recipe_config(). + is fetched to extract the inference container image URI. """ - - base_model_name = self._base_model_name() - if not base_model_name: - raise ValueError( - "base_model_name is required when deploying a model from an S3 checkpoint " - "(e.g. a serverful SMTJ training job) because no model package is available " - "to auto-resolve the inference container image. " - "Set trainer.base_model_name before calling build()." - ) - # If user already provided image_uri, skip hub resolution if self.image_uri: logger.info(f"Using provided image_uri: {self.image_uri}") return + base_model_name = self._base_model_name() hub_content_name = normalize_model_name(base_model_name) hub_name = getattr(self, "hub_name", None) or "SageMakerPublicHub" @@ -1338,7 +1328,6 @@ def _get_nova_hosting_config(self, instance_type=None): hub_content_name = model_package.inference_specification.containers[0].base_model.hub_content_name else: # No model package (e.g. SMTJ trainer): resolve from base_model_name - from sagemaker.train.common_utils.model_aliases import normalize_model_name base_model_name = self._base_model_name() hub_content_name = normalize_model_name(base_model_name) if base_model_name else None @@ -2985,7 +2974,16 @@ def _build_single_modelbuilder( self._fetch_and_cache_recipe_config() else: # No model package (e.g. serverful SMTJ training job). - # Resolve hosting config from Hub using base_model_name. + # base_model_name is required to identify the model type and resolve + # hosting config, escrow URI, tags, etc. + if not self._base_model_name(): + raise ValueError( + "trainer.base_model_name is required when deploying a model from an " + "S3 checkpoint (e.g. a serverful SMTJ training job) because no model " + "package is available to identify the model. " + "Set trainer.base_model_name before calling build()." + ) + # Resolve image_uri from Hub using base_model_name. self._resolve_hosting_config_from_base_model_name() # Nova models use a completely different deployment architecture From 7e88f21e9763bab28cff13afc397daaa93367300 Mon Sep 17 00:00:00 2001 From: jzhaoqwa Date: Tue, 28 Jul 2026 15:46:16 -0700 Subject: [PATCH 04/10] Improving logging level --- sagemaker-serve/src/sagemaker/serve/bedrock_model_builder.py | 4 ++-- sagemaker-serve/src/sagemaker/serve/model_builder.py | 4 ++-- .../src/sagemaker/train/common_utils/data_utils.py | 2 +- sagemaker-train/src/sagemaker/train/cpt_trainer.py | 1 - sagemaker-train/src/sagemaker/train/dpo_trainer.py | 1 - sagemaker-train/src/sagemaker/train/sft_trainer.py | 1 - 6 files changed, 5 insertions(+), 8 deletions(-) diff --git a/sagemaker-serve/src/sagemaker/serve/bedrock_model_builder.py b/sagemaker-serve/src/sagemaker/serve/bedrock_model_builder.py index e4a4625f4e..350cf28e9f 100644 --- a/sagemaker-serve/src/sagemaker/serve/bedrock_model_builder.py +++ b/sagemaker-serve/src/sagemaker/serve/bedrock_model_builder.py @@ -361,7 +361,7 @@ def deploy( self._get_bedrock_client(), model_arn ) if existing_deployment: - logger.warning( + logger.info( "Reusing existing custom model %s and deployment %s " "(matched model-source tag). No new resources were created. " "Pass reuse_resources=False to force new resources.", @@ -372,7 +372,7 @@ def deploy( "modelArn": model_arn, "customModelDeploymentArn": existing_deployment, } - logger.warning( + logger.info( "Reusing existing custom model %s (matched model-source tag); " "creating a new deployment on it. Pass reuse_resources=False to " "force a new model.", diff --git a/sagemaker-serve/src/sagemaker/serve/model_builder.py b/sagemaker-serve/src/sagemaker/serve/model_builder.py index b36ab609ab..67c51b713e 100644 --- a/sagemaker-serve/src/sagemaker/serve/model_builder.py +++ b/sagemaker-serve/src/sagemaker/serve/model_builder.py @@ -4056,7 +4056,7 @@ def build( reusable_endpoint = self._find_reusable_endpoint() if reusable_endpoint: self._reused_endpoint_name = reusable_endpoint - logger.warning( + logger.info( "Reusing existing Model %r (matched model-source tag). " "No new Model will be created. Pass reuse_resources=False " "to force a new Model.", @@ -5518,7 +5518,7 @@ def deploy( endpoint_name, reusable_endpoint, ) - logger.warning( + logger.info( "Reusing existing endpoint %r (matched model-source tag and " "deployment configuration). No new resources were created. " "Pass reuse_resources=False to force a new endpoint.", diff --git a/sagemaker-train/src/sagemaker/train/common_utils/data_utils.py b/sagemaker-train/src/sagemaker/train/common_utils/data_utils.py index f02c35084c..3c417b56e9 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/data_utils.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/data_utils.py @@ -248,7 +248,7 @@ def is_multimodal_data(dataset: Union[str, "DataSet"]) -> bool: True if multimodal fields detected, False otherwise """ - logger.info(f"Auto-detecting whether dataset is multimodal: {dataset}") + logger.debug(f"Auto-detecting whether dataset is multimodal: {dataset}") if isinstance(dataset, DataSet): data_s3_path = dataset.source diff --git a/sagemaker-train/src/sagemaker/train/cpt_trainer.py b/sagemaker-train/src/sagemaker/train/cpt_trainer.py index d151847005..51a0985a2d 100644 --- a/sagemaker-train/src/sagemaker/train/cpt_trainer.py +++ b/sagemaker-train/src/sagemaker/train/cpt_trainer.py @@ -40,7 +40,6 @@ from sagemaker.core.telemetry.constants import Feature logger = logging.getLogger(__name__) -logger.setLevel(logging.INFO) class CPTTrainer(BaseTrainer): diff --git a/sagemaker-train/src/sagemaker/train/dpo_trainer.py b/sagemaker-train/src/sagemaker/train/dpo_trainer.py index 47ce435705..bdfaf13884 100644 --- a/sagemaker-train/src/sagemaker/train/dpo_trainer.py +++ b/sagemaker-train/src/sagemaker/train/dpo_trainer.py @@ -31,7 +31,6 @@ from sagemaker.train.constants import get_sagemaker_hub_name logger = logging.getLogger(__name__) -logger.setLevel(logging.INFO) class DPOTrainer(BaseTrainer): diff --git a/sagemaker-train/src/sagemaker/train/sft_trainer.py b/sagemaker-train/src/sagemaker/train/sft_trainer.py index b8c147c9f2..4ef86b7c05 100644 --- a/sagemaker-train/src/sagemaker/train/sft_trainer.py +++ b/sagemaker-train/src/sagemaker/train/sft_trainer.py @@ -41,7 +41,6 @@ from sagemaker.core.training.constants import TrainingPlatform logger = logging.getLogger(__name__) -logger.setLevel(logging.INFO) class SFTTrainer(BaseTrainer): From 16b7cac4a276096632ad41dd227fbe338ebc8609 Mon Sep 17 00:00:00 2001 From: jzhaoqwa Date: Wed, 29 Jul 2026 15:58:59 -0700 Subject: [PATCH 05/10] fix: update job key used for show_metrics/stream_logs in MTRL --- .../src/sagemaker/train/base_trainer.py | 50 +++++++++++++------ .../train/common_utils/cloudwatch_metrics.py | 2 + .../sagemaker/train/multi_turn_rl_trainer.py | 2 + 3 files changed, 40 insertions(+), 14 deletions(-) diff --git a/sagemaker-train/src/sagemaker/train/base_trainer.py b/sagemaker-train/src/sagemaker/train/base_trainer.py index 468dacd610..0100c35005 100644 --- a/sagemaker-train/src/sagemaker/train/base_trainer.py +++ b/sagemaker-train/src/sagemaker/train/base_trainer.py @@ -396,11 +396,20 @@ def show_metrics( ValueError: If no training job has been run yet, no logs/metrics are found, or MLflow is not configured for OSS models. """ - # Validate that we have a training job to get metrics from - if not hasattr(self, '_latest_training_job') or self._latest_training_job is None: - raise ValueError( - "No training job found. Call .train() first, then call .show_metrics() " - "to view training metrics." + # Resolve the job reference. Prefer _latest_training_job (CreateTrainingJob), + # fall back to _latest_job (generic CreateJob API used by MTRL). + resolved_job = getattr(self, '_latest_training_job', None) + if resolved_job is None: + latest_job = getattr(self, '_latest_job', None) + if latest_job is None: + raise ValueError( + "No training job found. Call .train() first, then call .show_metrics() " + "to view training metrics. If training has already completed, set the " + "job name directly via trainer._latest_training_job = '' or " + "trainer._latest_job = ''." + ) + resolved_job = ( + latest_job.job_name if hasattr(latest_job, 'job_name') else str(latest_job) ) # Route based on model type @@ -408,18 +417,19 @@ def show_metrics( is_nova = _is_nova_model(model_name) if model_name else False if is_nova: - return self._show_metrics_cloudwatch(metrics, starting_step, ending_step, start_time, end_time) + return self._show_metrics_cloudwatch(resolved_job, metrics, starting_step, ending_step, start_time, end_time) else: - return self._show_metrics_mlflow(metrics, starting_step, ending_step) + return self._show_metrics_mlflow(resolved_job, metrics, starting_step, ending_step) def _show_metrics_mlflow( self, + resolved_job, metrics: Optional[List[str]] = None, starting_step: Optional[int] = None, ending_step: Optional[int] = None, ) -> None: """Pull and plot training metrics from MLflow for non-Nova models.""" - training_job = self._latest_training_job + training_job = resolved_job # Resolve the TrainingJob object if it's a string if isinstance(training_job, str): @@ -453,6 +463,7 @@ def _show_metrics_mlflow( def _show_metrics_cloudwatch( self, + resolved_job, metrics: Optional[List[str]] = None, starting_step: Optional[int] = None, ending_step: Optional[int] = None, @@ -461,7 +472,7 @@ def _show_metrics_cloudwatch( ) -> Any: """Parse and plot training metrics from CloudWatch logs (Nova models).""" - training_job = self._latest_training_job + training_job = resolved_job if hasattr(training_job, 'training_job_name'): job_id = training_job.training_job_name elif isinstance(training_job, str): @@ -628,10 +639,21 @@ def stream_logs(self, poll: int = 5, start_time: Optional[Any] = None) -> None: Raises: ValueError: If no training job has been run yet. """ - if not hasattr(self, '_latest_training_job') or self._latest_training_job is None: - raise ValueError( - "No training job found. Call .train(wait=False) first, " - "then call .stream_logs() to stream logs in real-time." + # Resolve the job reference. Prefer _latest_training_job (CreateTrainingJob), + # fall back to _latest_job (generic CreateJob API used by MTRL). + resolved_job = getattr(self, '_latest_training_job', None) + if resolved_job is None: + latest_job = getattr(self, '_latest_job', None) + if latest_job is None: + raise ValueError( + "No training job found. Call .train(wait=False) first, " + "then call .stream_logs() to stream logs in real-time. " + "If training has already completed, set the job name directly via " + "trainer._latest_training_job = '' or " + "trainer._latest_job = ''." + ) + resolved_job = ( + latest_job.job_name if hasattr(latest_job, 'job_name') else str(latest_job) ) # Resolve start_time for SMHP jobs @@ -642,7 +664,7 @@ def stream_logs(self, poll: int = 5, start_time: Optional[Any] = None) -> None: else: start_time_ms = int(start_time) - training_job = self._latest_training_job + training_job = resolved_job compute = getattr(self, 'compute', None) if isinstance(compute, HyperPodCompute): diff --git a/sagemaker-train/src/sagemaker/train/common_utils/cloudwatch_metrics.py b/sagemaker-train/src/sagemaker/train/common_utils/cloudwatch_metrics.py index ed1906675f..7686a763ee 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/cloudwatch_metrics.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/cloudwatch_metrics.py @@ -33,11 +33,13 @@ "SFT": {"training_loss": TRAINING_LOSS_REGEX, "lr": LEARNING_RATE_REGEX}, "CPT": {"training_loss": TRAINING_LOSS_REGEX, "lr": LEARNING_RATE_REGEX}, "RLVR": {"reward_score": SMTJ_RLVR_REWARD_SCORE_REGEX}, + "MTRL": {"reward_score": SMTJ_RLVR_REWARD_SCORE_REGEX}, }, "smhp": { "SFT": {"training_loss": TRAINING_LOSS_REGEX, "lr": LEARNING_RATE_REGEX}, "CPT": {"training_loss": TRAINING_LOSS_REGEX, "lr": LEARNING_RATE_REGEX}, "RLVR": {"reward_score": SMHP_RLVR_REWARD_SCORE_REGEX}, + "MTRL": {"reward_score": SMHP_RLVR_REWARD_SCORE_REGEX}, }, } diff --git a/sagemaker-train/src/sagemaker/train/multi_turn_rl_trainer.py b/sagemaker-train/src/sagemaker/train/multi_turn_rl_trainer.py index 507d38202a..305c268a44 100644 --- a/sagemaker-train/src/sagemaker/train/multi_turn_rl_trainer.py +++ b/sagemaker-train/src/sagemaker/train/multi_turn_rl_trainer.py @@ -173,6 +173,8 @@ class MultiTurnRLTrainer(BaseTrainer): and 'job_name_prefix'. If not specified, no notifications are sent. """ + _customization_technique = "MTRL" + def __init__( self, model: Union[str, ModelPackage], From 5156fa38596c3a0feca4c6aeb6bc0ff7fd684e38 Mon Sep 17 00:00:00 2001 From: jzhaoqwa Date: Thu, 30 Jul 2026 15:10:54 -0700 Subject: [PATCH 06/10] fix(serve): Speed up reuse_resources with Tagging API and resolve string training jobs Use resourcegroupstaggingapi.get_resources() for O(1) tag lookups instead of scanning all models/endpoints, and auto-resolve string _latest_training_job to TrainingJob objects in ModelBuilder so trainers work without manual .get(). --- .../src/sagemaker/serve/model_builder.py | 45 +++---- .../src/sagemaker/serve/model_reuse.py | 118 +++++++++++++++++- 2 files changed, 140 insertions(+), 23 deletions(-) diff --git a/sagemaker-serve/src/sagemaker/serve/model_builder.py b/sagemaker-serve/src/sagemaker/serve/model_builder.py index c125a749f4..5e9903cd7a 100644 --- a/sagemaker-serve/src/sagemaker/serve/model_builder.py +++ b/sagemaker-serve/src/sagemaker/serve/model_builder.py @@ -1811,6 +1811,11 @@ def _is_model_customization(self) -> bool: # ModelTrainer with model customization if isinstance(self.model, ModelTrainer) and hasattr(self.model, "_latest_training_job"): + # Resolve string job name to TrainingJob object if needed + if isinstance(self.model._latest_training_job, str): + self.model._latest_training_job = TrainingJob.get( + training_job_name=self.model._latest_training_job + ) # Check model_package_config first (new location) if ( hasattr(self.model._latest_training_job, "model_package_config") @@ -1845,6 +1850,11 @@ def _is_model_customization(self) -> bool: if isinstance(self.model, MultiTurnRLTrainer): return True if isinstance(self.model, BaseTrainer) and hasattr(self.model, "_latest_training_job"): + # Resolve string job name to TrainingJob object if needed + if isinstance(self.model._latest_training_job, str): + self.model._latest_training_job = TrainingJob.get( + training_job_name=self.model._latest_training_job + ) # Trainer built from an S3 checkpoint (e.g. Serverful SMTJ): no model # package, but the completed training job has an S3 output path that # holds the customized artifacts. @@ -1926,6 +1936,11 @@ def _fetch_model_package_arn(self) -> Optional[str]: return arn if hasattr(self.model, "_latest_training_job"): + # Resolve string job name to TrainingJob object if needed + if isinstance(self.model._latest_training_job, str): + self.model._latest_training_job = TrainingJob.get( + training_job_name=self.model._latest_training_job + ) # Try output_model_package_arn first (preferred) if hasattr(self.model._latest_training_job, "output_model_package_arn"): arn = self.model._latest_training_job.output_model_package_arn @@ -2049,38 +2064,24 @@ def _get_model_for_endpoint(self, endpoint_name: str) -> Optional[Model]: def _find_reusable_model(self) -> Optional["Model"]: """Find an existing SageMaker Model tagged with the same model source. - Scans Models by creation time and checks for a matching model-source tag. + Uses the Resource Groups Tagging API for efficient server-side tag + filtering when available, falling back to paginated list+list_tags scan. Returns the Model resource if found (and it still exists), None otherwise. """ source_id = self._resolve_model_source_id() if not source_id: return None - from sagemaker.serve.model_reuse import normalize_tag_value + from sagemaker.serve.model_reuse import normalize_tag_value, find_sagemaker_model_arn_by_tag tag_value = normalize_tag_value(source_id) sagemaker_client = self.sagemaker_session.sagemaker_client try: - next_token = None - while True: - kwargs = {"SortBy": "CreationTime", "SortOrder": "Descending"} - if next_token: - kwargs["NextToken"] = next_token - response = sagemaker_client.list_models(**kwargs) - for model_summary in response.get("Models", []): - model_name = model_summary.get("ModelName") - model_arn = model_summary.get("ModelArn") - if not model_arn: - continue - tags = sagemaker_client.list_tags(ResourceArn=model_arn).get("Tags", []) - if any( - t.get("Key") == MODEL_SOURCE_TAG_KEY and t.get("Value") == tag_value - for t in tags - ): - return Model.get(model_name=model_name, region=self.region) - next_token = response.get("NextToken") - if not next_token: - break + model_arn = find_sagemaker_model_arn_by_tag(sagemaker_client, tag_value) + if model_arn: + # Extract model name from ARN: arn:aws:sagemaker:region:account:model/name + model_name = model_arn.rsplit("/", 1)[-1] + return Model.get(model_name=model_name, region=self.region) except Exception as e: logger.warning("Could not search Models for reuse: %s", e) diff --git a/sagemaker-serve/src/sagemaker/serve/model_reuse.py b/sagemaker-serve/src/sagemaker/serve/model_reuse.py index b534173d1f..06bfef5866 100644 --- a/sagemaker-serve/src/sagemaker/serve/model_reuse.py +++ b/sagemaker-serve/src/sagemaker/serve/model_reuse.py @@ -218,8 +218,124 @@ def _bedrock_resource_has_tag(bedrock_client, resource_arn: str, tag_value: str) ) +def _find_resource_arn_by_tagging_api( + sagemaker_client, tag_value: str, resource_type: str +) -> Optional[str]: + """Use Resource Groups Tagging API to find a resource by tag (fast path). + + Makes a single server-side filtered query instead of iterating through all + resources and calling list_tags on each one. + + Args: + sagemaker_client: A boto3 SageMaker client (used to derive the session). + tag_value: The normalized tag value to search for. + resource_type: The resource type filter (e.g. "sagemaker:model", + "sagemaker:endpoint"). + + Returns: + The resource ARN if found, empty string "" if no match (signals fast path + completed successfully with no results), or None if the tagging API is + unavailable (signals caller should fall back to the slow scan). + """ + try: + # Extract region from the SageMaker client + import boto3 as _boto3 + + region = sagemaker_client.meta.region_name + if not region: + return None + + tagging_client = _boto3.client( + "resourcegroupstaggingapi", + region_name=region, + ) + + pagination_token = "" + while True: + kwargs = { + "TagFilters": [ + {"Key": MODEL_SOURCE_TAG_KEY, "Values": [tag_value]} + ], + "ResourceTypeFilters": [resource_type], + } + if pagination_token: + kwargs["PaginationToken"] = pagination_token + + response = tagging_client.get_resources(**kwargs) + for mapping in response.get("ResourceTagMappingList", []): + arn = mapping.get("ResourceARN") + if arn: + return arn + + pagination_token = response.get("PaginationToken", "") + if not pagination_token: + return "" # Fast path succeeded, no matching resource found + + except ClientError as e: + error_code = e.response.get("Error", {}).get("Code", "") + if error_code == _ACCESS_DENIED_CODE: + logger.debug( + "Resource Groups Tagging API access denied (tag:GetResources). " + "Falling back to paginated list+list_tags scan." + ) + return None # Signal caller to use fallback + logger.debug("Resource Groups Tagging API call failed: %s. Using fallback.", e) + return None + except Exception as e: + logger.debug("Resource Groups Tagging API unavailable: %s. Using fallback.", e) + return None + + +def find_sagemaker_model_arn_by_tag(sagemaker_client, tag_value: str) -> Optional[str]: + """Return the ARN of the first SageMaker Model carrying the source tag. + + Uses the Resource Groups Tagging API for efficient server-side filtering + when available, falling back to paginated list+list_tags if denied. + + Args: + sagemaker_client: A boto3 SageMaker client. + tag_value: The normalized tag value to match. + + Returns: + Model ARN if found, None otherwise. + """ + # Try the fast path first: Resource Groups Tagging API + arn = _find_resource_arn_by_tagging_api( + sagemaker_client, tag_value, resource_type="sagemaker:model" + ) + if arn is not None: + return arn if arn != "" else None + + # Fallback: paginate through all models and check tags individually + next_token = None + while True: + kwargs = {"SortBy": "CreationTime", "SortOrder": "Descending"} + if next_token: + kwargs["NextToken"] = next_token + response = sagemaker_client.list_models(**kwargs) + for model_summary in response.get("Models", []): + model_arn = model_summary.get("ModelArn") + if model_arn and _sagemaker_resource_has_tag(sagemaker_client, model_arn, tag_value): + return model_arn + next_token = response.get("NextToken") + if not next_token: + return None + + def _find_sagemaker_endpoint_arn_by_tag(sagemaker_client, tag_value: str) -> Optional[str]: - """Return the ARN of the first SageMaker endpoint carrying the source tag.""" + """Return the ARN of the first SageMaker endpoint carrying the source tag. + + Uses the Resource Groups Tagging API for efficient server-side filtering + when available, falling back to paginated list+list_tags if denied. + """ + # Try the fast path first: Resource Groups Tagging API + arn = _find_resource_arn_by_tagging_api( + sagemaker_client, tag_value, resource_type="sagemaker:endpoint" + ) + if arn is not None: + return arn if arn != "" else None + + # Fallback: paginate through all endpoints and check tags individually next_token = None while True: kwargs = {"NextToken": next_token} if next_token else {} From 3133f26e1a115f518a937c1d39c4e3b949823835 Mon Sep 17 00:00:00 2001 From: jzhaoqwa Date: Thu, 30 Jul 2026 15:45:16 -0700 Subject: [PATCH 07/10] Remove region logging --- sagemaker-core/src/sagemaker/core/utils/utils.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/sagemaker-core/src/sagemaker/core/utils/utils.py b/sagemaker-core/src/sagemaker/core/utils/utils.py index 865ae0f96e..8874d5fb58 100644 --- a/sagemaker-core/src/sagemaker/core/utils/utils.py +++ b/sagemaker-core/src/sagemaker/core/utils/utils.py @@ -369,10 +369,8 @@ def __init__( self.session = session self.region_name = region_name # Read region from environment variable, default to us-west-2 - import os env_region = os.environ.get('SAGEMAKER_REGION', region_name) env_stage = os.environ.get('SAGEMAKER_STAGE', 'prod') # default to gamma - logger.info(f"Runs on sagemaker {env_stage}, region:{env_region}") endpoint_url = os.environ.get('SAGEMAKER_ENDPOINT') runtime_endpoint_url = os.environ.get('SAGEMAKER_RUNTIME_ENDPOINT') From 1d5ad47df88a3996a13bb5c24ea493996760ca92 Mon Sep 17 00:00:00 2001 From: jzhaoqwa Date: Thu, 30 Jul 2026 16:57:26 -0700 Subject: [PATCH 08/10] fix(train): Improve trainer UX and reduce verbose logging - Cache sagemaker_session in BaseTrainer.__init__ to avoid creating duplicate sessions on every method call - Remove redundant role validation (was validating 3x per train() call), now validates once in ModelTrainer.__init__ - Demote noisy INFO logs to DEBUG (role validated, stopping condition defaults, recipe paths, output compression) - Add num_lines param to stream_logs() to limit output for long jobs - Prefix CloudWatch log lines with [CloudWatch] and use print() to distinguish container output from SDK logging - Improve show_metrics() error message when time range yields no logs - Raise ValueError on AccessDenied in dry_run data path validation instead of silently warning - Improve model_package_group error message to mention compute option - Move local imports in _train_serverful_smtj to top-level - Remove redundant get_role() call in ModelTrainer.from_recipe() --- .../core/helper/iam_role_resolver.py | 2 +- .../src/sagemaker/train/base_trainer.py | 60 +++++++++++-------- .../train/common_utils/cloudwatch_metrics.py | 10 +++- .../train/common_utils/data_utils.py | 10 ++-- .../train/common_utils/finetune_utils.py | 8 ++- .../src/sagemaker/train/defaults.py | 6 +- .../src/sagemaker/train/model_trainer.py | 1 - 7 files changed, 58 insertions(+), 39 deletions(-) diff --git a/sagemaker-core/src/sagemaker/core/helper/iam_role_resolver.py b/sagemaker-core/src/sagemaker/core/helper/iam_role_resolver.py index db4023b573..3baa1205da 100644 --- a/sagemaker-core/src/sagemaker/core/helper/iam_role_resolver.py +++ b/sagemaker-core/src/sagemaker/core/helper/iam_role_resolver.py @@ -601,7 +601,7 @@ def resolve_and_validate_role( role_type, ) else: - logger.info("Role '%s' validated for %s. Using it.", role_arn, role_type) + logger.debug("Role '%s' validated for %s. Using it.", role_arn, role_type) return role_arn diff --git a/sagemaker-train/src/sagemaker/train/base_trainer.py b/sagemaker-train/src/sagemaker/train/base_trainer.py index c367fb1154..f175d33013 100644 --- a/sagemaker-train/src/sagemaker/train/base_trainer.py +++ b/sagemaker-train/src/sagemaker/train/base_trainer.py @@ -17,9 +17,10 @@ import boto3 from sagemaker.core.helper.session_helper import Session -from sagemaker.core.training.configs import Tag, Networking, InputData, Channel, OutputDataConfig, HyperPodCompute +from sagemaker.core.training.configs import Tag, Networking, InputData, Channel, OutputDataConfig, HyperPodCompute, TrainingJobCompute from sagemaker.core.utils.logs import MultiLogStreamHandler from sagemaker.core.shapes import shapes +from sagemaker.core.shapes import S3DataSource from sagemaker.core.resources import TrainingJob from sagemaker.train.common_utils.recipe_utils import _is_nova_model, resolve_recipe, get_resolved_recipe_from_context, NoRecipeError from sagemaker.core.s3.utils import resolve_s3_uri_placeholders @@ -40,6 +41,7 @@ from sagemaker.train.common_utils.validator import validate_hyperpod_compute from sagemaker.train.common_utils.cloudwatch_metrics import fetch_and_plot_metrics, _get_smhp_log_group from sagemaker.train.defaults import TrainDefaults +from sagemaker.train.model_trainer import ModelTrainer from sagemaker.train.utils import _get_unique_name logger = logging.getLogger(__name__) @@ -113,7 +115,7 @@ def __init__( disable_output_compression: Optional[bool] = False, notifications: Optional[Dict[str, Any]] = None, ): - self.sagemaker_session = sagemaker_session + self.sagemaker_session = sagemaker_session or TrainDefaults.get_sagemaker_session() self.role = role self.base_job_name = base_job_name self.tags = tags @@ -624,7 +626,7 @@ def list_notification_rules( event_bus_arn=event_bus_arn, ) - def stream_logs(self, poll: int = 5, start_time: Optional[Any] = None) -> None: + def stream_logs(self, poll: int = 5, start_time: Optional[Any] = None, num_lines: Optional[int] = None) -> None: """Stream CloudWatch logs in real-time (like ``kubectl logs -f``). Continuously polls for new log events and prints them as they arrive. @@ -638,6 +640,10 @@ def stream_logs(self, poll: int = 5, start_time: Optional[Any] = None) -> None: attaching to a job that's already running. If not provided, auto-resolved from the training job's start time (SMTJ) or defaults to now (HyperPod). + num_lines: Optional maximum number of log lines to print. When + specified, streaming stops after this many lines have been + printed. Useful for long jobs where the full log is too verbose. + If not provided, streams all logs until the job completes. Raises: ValueError: If no training job has been run yet. @@ -671,11 +677,11 @@ def stream_logs(self, poll: int = 5, start_time: Optional[Any] = None) -> None: compute = getattr(self, 'compute', None) if isinstance(compute, HyperPodCompute): - self._stream_logs_smhp(training_job, compute, poll, start_time_ms) + self._stream_logs_smhp(training_job, compute, poll, start_time_ms, num_lines=num_lines) else: - self._stream_logs_smtj(training_job, poll) + self._stream_logs_smtj(training_job, poll, num_lines=num_lines) - def _stream_logs_smtj(self, training_job, poll: int) -> None: + def _stream_logs_smtj(self, training_job, poll: int, num_lines: Optional[int] = None) -> None: """Stream logs for an SMTJ training job using MultiLogStreamHandler.""" # Resolve job name @@ -699,12 +705,18 @@ def _stream_logs_smtj(self, training_job, poll: int) -> None: logger.info(f"Log group: {log_group}") terminal_statuses = {"Completed", "Failed", "Stopped"} + lines_printed = 0 + _CW_PREFIX = "[CloudWatch] " while True: for stream_name, event in handler.get_latest_log_events(): message = event.get("message", "").rstrip() if message: - logger.info(message) + print(f"{_CW_PREFIX}{message}") + lines_printed += 1 + if num_lines and lines_printed >= num_lines: + logger.info(f"Reached num_lines limit ({num_lines}). Stopping log stream.") + return # Check job status try: @@ -715,7 +727,11 @@ def _stream_logs_smtj(self, training_job, poll: int) -> None: for stream_name, event in handler.get_latest_log_events(): message = event.get("message", "").rstrip() if message: - logger.info(message) + print(f"{_CW_PREFIX}{message}") + lines_printed += 1 + if num_lines and lines_printed >= num_lines: + logger.info(f"Reached num_lines limit ({num_lines}). Stopping log stream.") + return logger.info(f"Job {job_name} finished with status: {status}") return except Exception: @@ -723,7 +739,7 @@ def _stream_logs_smtj(self, training_job, poll: int) -> None: time.sleep(poll) - def _stream_logs_smhp(self, training_job, compute, poll: int, start_time_ms=None) -> None: + def _stream_logs_smhp(self, training_job, compute, poll: int, start_time_ms=None, num_lines: Optional[int] = None) -> None: """Stream logs for a HyperPod job using filter_log_events polling.""" if isinstance(training_job, str): @@ -756,6 +772,8 @@ def _stream_logs_smhp(self, training_job, compute, poll: int, start_time_ms=None else: last_timestamp = int(time.time() * 1000) seen_event_ids = set() + lines_printed = 0 + _CW_PREFIX = "[CloudWatch] " while True: try: @@ -774,7 +792,11 @@ def _stream_logs_smhp(self, training_job, compute, poll: int, start_time_ms=None seen_event_ids.add(event_id) message = event.get("message", "").rstrip() if message: - logger.info(message) + print(f"{_CW_PREFIX}{message}") + lines_printed += 1 + if num_lines and lines_printed >= num_lines: + logger.info(f"Reached num_lines limit ({num_lines}). Stopping log stream.") + return ts = event.get("timestamp", 0) if ts > last_timestamp: last_timestamp = ts @@ -848,22 +870,10 @@ def _train_serverful_smtj(self, training_dataset=None, validation_dataset=None, from ``self._customization_technique``) and any extra hyperparameters from ``_get_extra_smtj_hyperparameters()``. """ - import logging - import tempfile - from sagemaker.train.model_trainer import ModelTrainer - from sagemaker.core.training.configs import TrainingJobCompute, InputData, Networking - from sagemaker.core.shapes import S3DataSource - from sagemaker.train.common_utils.finetune_utils import ( - get_recipe_s3_uri, - get_training_image, - _validate_hyperparameter_values, - ) - from sagemaker.train.defaults import TrainDefaults - sagemaker_session = TrainDefaults.get_sagemaker_session( sagemaker_session=self.sagemaker_session ) - role = TrainDefaults.get_role(role=self.role, sagemaker_session=sagemaker_session) + role = self.role compute = self.compute customization_technique = self._customization_technique @@ -876,7 +886,7 @@ def _train_serverful_smtj(self, training_dataset=None, validation_dataset=None, sagemaker_session=sagemaker_session, ) - logger.info(f"SMTJ recipe S3 URI: {recipe_s3_uri}") + logger.debug(f"SMTJ recipe S3 URI: {recipe_s3_uri}") # Download recipe from S3 to a local temp file recipe_s3_uri = resolve_s3_uri_placeholders(recipe_s3_uri, sagemaker_session) @@ -1095,7 +1105,7 @@ def _yaml_safe_default(value): with open(recipe_local_path, "w") as f: f.write(recipe_content) - logger.info(f"Recipe downloaded and rendered to: {recipe_local_path}") + logger.debug(f"Recipe downloaded and rendered to: {recipe_local_path}") # Resolve training image training_image = self.training_image diff --git a/sagemaker-train/src/sagemaker/train/common_utils/cloudwatch_metrics.py b/sagemaker-train/src/sagemaker/train/common_utils/cloudwatch_metrics.py index 4bec69559a..fa7fc30a85 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/cloudwatch_metrics.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/cloudwatch_metrics.py @@ -445,9 +445,17 @@ def fetch_and_plot_metrics( ) if not log_events: + time_hint = "" + if start_time or end_time: + time_hint = ( + f" No logs were found in the specified time range " + f"(start_time={start_time}, end_time={end_time}). " + f"Try adjusting the time range or omitting start_time/end_time " + f"to search the full job duration." + ) raise ValueError( f"No CloudWatch logs found for job '{job_id}' in log group '{log_group}'. " - f"The job may still be starting, or logs may not be available yet." + f"The job may still be starting, or logs may not be available yet.{time_hint}" ) # Parse metrics from logs diff --git a/sagemaker-train/src/sagemaker/train/common_utils/data_utils.py b/sagemaker-train/src/sagemaker/train/common_utils/data_utils.py index 3c417b56e9..7a2b7dda8b 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/data_utils.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/data_utils.py @@ -140,12 +140,10 @@ def validate_data_path_exists( except ClientError as e: code = e.response["Error"]["Code"] if code == "403" or "AccessDenied" in str(e): - # Caller may not have access but the execution role might — - # log a warning and allow the job to proceed. - logger.warning( - "Cannot verify S3 %s path %s from caller identity " - "(AccessDenied). The execution role may still have access.", - label, data_path, + raise ValueError( + f"Cannot verify S3 {label} path '{data_path}': access denied. " + f"Ensure the path exists and the caller has s3:ListBucket and " + f"s3:GetObject permissions on the bucket." ) else: raise ValueError(f"Error accessing S3 {label} path {data_path}: {e}") 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 94f5f7af3e..a3bd1fba97 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/finetune_utils.py @@ -1123,8 +1123,12 @@ def _validate_and_resolve_model_package_group(model, model_package_group_name): if isinstance(model, ModelPackage): return model.model_package_group_name - raise ValueError("model_package_group_name must be provided when model given is " - "not a ModelPackage artifact/not continued finetuning") + raise ValueError( + "model_package_group is required for serverless training (when compute is not set). " + "Either provide model_package_group to store the fine-tuned model, or set " + "compute=TrainingJobCompute(...) / HyperPodCompute(...) to use managed compute " + "where model_package_group is optional." + ) def _validate_eula_for_gated_model(model, accept_eula, is_gated_model): diff --git a/sagemaker-train/src/sagemaker/train/defaults.py b/sagemaker-train/src/sagemaker/train/defaults.py index c358c11a04..dbd5cdfb0d 100644 --- a/sagemaker-train/src/sagemaker/train/defaults.py +++ b/sagemaker-train/src/sagemaker/train/defaults.py @@ -164,10 +164,10 @@ def get_stopping_condition( max_pending_time_in_seconds=None, max_wait_time_in_seconds=None, ) - logger.info(f"StoppingCondition not provided. Using default:\n{stopping_condition}") + logger.debug(f"StoppingCondition not provided. Using default:\n{stopping_condition}") if stopping_condition.max_runtime_in_seconds is None: stopping_condition.max_runtime_in_seconds = DEFAULT_MAX_RUNTIME_IN_SECONDS - logger.info( + logger.debug( "Max runtime not provided. Using default:\n" f"{stopping_condition.max_runtime_in_seconds}" ) @@ -201,7 +201,7 @@ def get_output_data_config( ) if output_data_config.compression_type is None: output_data_config.compression_type = "GZIP" - logger.info( + logger.debug( f"OutputDataConfig compression type not provided. Using default:\n" f"{output_data_config.compression_type}" ) diff --git a/sagemaker-train/src/sagemaker/train/model_trainer.py b/sagemaker-train/src/sagemaker/train/model_trainer.py index 3554a74ea7..15aa4f0d6b 100644 --- a/sagemaker-train/src/sagemaker/train/model_trainer.py +++ b/sagemaker-train/src/sagemaker/train/model_trainer.py @@ -1283,7 +1283,6 @@ def from_recipe( raise ValueError("training_image must be provided when using training_image_config.") sagemaker_session = TrainDefaults.get_sagemaker_session(sagemaker_session) - role = TrainDefaults.get_role(role=role, sagemaker_session=sagemaker_session) # The training recipe is used to prepare the following args: # - source_code From 1f055259bf8664b5af00ad649e10c98a9a4828c7 Mon Sep 17 00:00:00 2001 From: jzhaoqwa Date: Thu, 30 Jul 2026 17:04:41 -0700 Subject: [PATCH 09/10] resolve conflict --- .../src/sagemaker/train/base_trainer.py | 6 ++--- .../train/common_utils/log_streamer.py | 25 ++++++++++++++++--- 2 files changed, 25 insertions(+), 6 deletions(-) diff --git a/sagemaker-train/src/sagemaker/train/base_trainer.py b/sagemaker-train/src/sagemaker/train/base_trainer.py index 872751bf68..2d36016256 100644 --- a/sagemaker-train/src/sagemaker/train/base_trainer.py +++ b/sagemaker-train/src/sagemaker/train/base_trainer.py @@ -680,9 +680,9 @@ def stream_logs(self, poll: int = 5, start_time: Optional[Any] = None, num_lines if isinstance(compute, HyperPodCompute): self._stream_logs_smhp(training_job, compute, poll, start_time_ms, num_lines=num_lines) else: - self._stream_logs_smtj(training_job, poll, start_time_ms) + self._stream_logs_smtj(training_job, poll, start_time_ms, num_lines=num_lines) - def _stream_logs_smtj(self, training_job, poll: int, start_time_ms=None) -> None: + def _stream_logs_smtj(self, training_job, poll: int, start_time_ms=None, num_lines: Optional[int] = None) -> None: """Stream logs for an SMTJ training job.""" from sagemaker.train.common_utils.log_streamer import ( LogStreamer, @@ -714,7 +714,7 @@ def _get_status() -> str: job = TrainingJob.get(training_job_name=job_name) return job.training_job_status - stream_log_loop(streamer, poll, _get_status) + stream_log_loop(streamer, poll, _get_status, num_lines=num_lines) def _stream_logs_smhp(self, training_job, compute, poll: int, start_time_ms=None, num_lines: Optional[int] = None) -> None: """Stream logs for a HyperPod job using filter_log_events polling.""" diff --git a/sagemaker-train/src/sagemaker/train/common_utils/log_streamer.py b/sagemaker-train/src/sagemaker/train/common_utils/log_streamer.py index c120bd9131..3538c95200 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/log_streamer.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/log_streamer.py @@ -227,6 +227,7 @@ def stream_log_loop( streamer: LogStreamer, poll: int, status_fn: Callable[[], str], + num_lines: Optional[int] = None, ) -> None: """Run the standard log streaming loop. @@ -237,7 +238,22 @@ def stream_log_loop( :param streamer: A configured LogStreamer instance. :param poll: Seconds between polls. :param status_fn: Callable that returns the current job status string. + :param num_lines: Optional maximum number of log lines to print. + When specified, streaming stops after this many lines. """ + _CW_PREFIX = "[CloudWatch] " + lines_printed = 0 + + def _print_event(ts_ms: int, message: str) -> bool: + """Print a log event. Returns True if num_lines limit reached.""" + nonlocal lines_printed + print(f"{_CW_PREFIX}[{_format_timestamp(ts_ms)}] {message}") + lines_printed += 1 + if num_lines and lines_printed >= num_lines: + logger.info("Reached num_lines limit (%d). Stopping log stream.", num_lines) + return True + return False + status = status_fn() if status in TERMINAL_STATUSES: logger.info("Job already in terminal state: %s", status) @@ -247,7 +263,8 @@ def stream_log_loop( if not events: break for ts_ms, message in events: - logger.info("[%s] %s", _format_timestamp(ts_ms), message) + if _print_event(ts_ms, message): + return except ClientError: pass logger.info("Job finished with status: %s", status) @@ -285,7 +302,8 @@ def stream_log_loop( if events: empty_cycles = 0 for ts_ms, message in events: - logger.info("[%s] %s", _format_timestamp(ts_ms), message) + if _print_event(ts_ms, message): + return else: empty_cycles += 1 if empty_cycles == warn_cycle: @@ -297,7 +315,8 @@ def stream_log_loop( status = status_fn() if status in TERMINAL_STATUSES: for ts_ms, message in streamer.poll_once(): - logger.info("[%s] %s", _format_timestamp(ts_ms), message) + if _print_event(ts_ms, message): + return logger.info("Job finished with status: %s", status) return From 17bd275adae5cef4f1e29f3bd75795823204cac5 Mon Sep 17 00:00:00 2001 From: jzhaoqwa Date: Fri, 31 Jul 2026 12:26:31 -0700 Subject: [PATCH 10/10] Update import --- sagemaker-serve/src/sagemaker/serve/model_reuse.py | 13 +++++-------- 1 file changed, 5 insertions(+), 8 deletions(-) diff --git a/sagemaker-serve/src/sagemaker/serve/model_reuse.py b/sagemaker-serve/src/sagemaker/serve/model_reuse.py index 06bfef5866..f08dbce3a0 100644 --- a/sagemaker-serve/src/sagemaker/serve/model_reuse.py +++ b/sagemaker-serve/src/sagemaker/serve/model_reuse.py @@ -20,6 +20,8 @@ from botocore.exceptions import ClientError +import boto3 + logger = logging.getLogger(__name__) MODEL_SOURCE_TAG_KEY = "sagemaker.amazonaws.com/model-source" @@ -227,7 +229,8 @@ def _find_resource_arn_by_tagging_api( resources and calling list_tags on each one. Args: - sagemaker_client: A boto3 SageMaker client (used to derive the session). + sagemaker_client: A boto3 SageMaker client (used for region detection + and as fallback for creating the tagging client). tag_value: The normalized tag value to search for. resource_type: The resource type filter (e.g. "sagemaker:model", "sagemaker:endpoint"). @@ -238,17 +241,11 @@ def _find_resource_arn_by_tagging_api( unavailable (signals caller should fall back to the slow scan). """ try: - # Extract region from the SageMaker client - import boto3 as _boto3 - region = sagemaker_client.meta.region_name if not region: return None - tagging_client = _boto3.client( - "resourcegroupstaggingapi", - region_name=region, - ) + tagging_client = boto3.client("resourcegroupstaggingapi", region_name=region) pagination_token = "" while True: