diff --git a/sagemaker-train/src/sagemaker/train/model_trainer.py b/sagemaker-train/src/sagemaker/train/model_trainer.py index 35b1395cce..8a054a1de3 100644 --- a/sagemaker-train/src/sagemaker/train/model_trainer.py +++ b/sagemaker-train/src/sagemaker/train/model_trainer.py @@ -343,6 +343,57 @@ class ModelTrainer(BaseModel): config_mgr: SageMakerConfig = SageMakerConfig() + @property + def output_data(self) -> str: + """The S3 URI of the training job's ``output.tar.gz`` archive. + + This is the non-model output archive that SageMaker uploads from the + container's ``/opt/ml/output/data`` directory after training completes, + distinct from the model artifact (``model.tar.gz``). The URI follows the + layout ``{s3_output_path}/{training_job_name}/output/output.tar.gz``. + + The S3 output path is resolved from the latest training job's + ``output_data_config`` and falls back to the trainer's own + ``output_data_config`` when the job does not report one. + + Note: + This is a computed S3 URI derived from the job's configuration; it is + not verified to exist. The archive is only present once the job has + completed successfully and produced output data. + + Returns: + str: The fully-qualified S3 URI to the job's ``output.tar.gz``. + + Raises: + ValueError: If no training job has been created yet (``train`` has + not been called), or if no S3 output path can be resolved. + """ + training_job = self._latest_training_job + if training_job is None: + raise ValueError( + "No training job is associated with this ModelTrainer. " + "Call train() before accessing output_data." + ) + + # Prefer the S3 output path reported by the job, falling back to the + # trainer's own config. ``Unassigned`` and ``None`` are both falsy. + s3_output_path = None + job_output_config = getattr(training_job, "output_data_config", None) + if job_output_config: + s3_output_path = getattr(job_output_config, "s3_output_path", None) + if not s3_output_path and self.output_data_config is not None: + s3_output_path = getattr(self.output_data_config, "s3_output_path", None) + + if not s3_output_path: + raise ValueError( + "Unable to resolve the output S3 path for this training job. " + "Ensure output_data_config is set on the ModelTrainer or the " + "training job." + ) + + s3_output_path = s3_output_path.rstrip("/") + return f"{s3_output_path}/{training_job.training_job_name}/output/output.tar.gz" + def _populate_intelligent_defaults(self): """Function to populate all the possible default configs diff --git a/sagemaker-train/tests/unit/train/test_model_trainer.py b/sagemaker-train/tests/unit/train/test_model_trainer.py index ceb371dfa6..0b5f9e3b0a 100644 --- a/sagemaker-train/tests/unit/train/test_model_trainer.py +++ b/sagemaker-train/tests/unit/train/test_model_trainer.py @@ -2220,3 +2220,71 @@ def test_log_actionable_client_error_other_codes_stay_silent(caplog): _log_actionable_client_error(error) assert caplog.text == "" + + +def test_output_data_returns_output_tar_gz_uri(model_trainer): + """output_data derives the output.tar.gz S3 URI from the completed job.""" + from sagemaker.core.shapes import OutputDataConfig as CoreOutputDataConfig + + model_trainer._latest_training_job = TrainingJob( + training_job_name="my-training-job", + output_data_config=CoreOutputDataConfig( + s3_output_path=f"s3://{DEFAULT_BUCKET}/{DEFAULT_BUCKET_PREFIX}" + ), + ) + + assert model_trainer.output_data == ( + f"s3://{DEFAULT_BUCKET}/{DEFAULT_BUCKET_PREFIX}" "/my-training-job/output/output.tar.gz" + ) + + +def test_output_data_strips_trailing_slash(model_trainer): + """A trailing slash on s3_output_path must not produce a double slash.""" + from sagemaker.core.shapes import OutputDataConfig as CoreOutputDataConfig + + model_trainer._latest_training_job = TrainingJob( + training_job_name="my-training-job", + output_data_config=CoreOutputDataConfig( + s3_output_path=f"s3://{DEFAULT_BUCKET}/{DEFAULT_BUCKET_PREFIX}/" + ), + ) + + assert model_trainer.output_data == ( + f"s3://{DEFAULT_BUCKET}/{DEFAULT_BUCKET_PREFIX}" "/my-training-job/output/output.tar.gz" + ) + + +def test_output_data_falls_back_to_trainer_output_config(model_trainer): + """When the job resource has no output_data_config, fall back to the trainer's.""" + model_trainer._latest_training_job = TrainingJob(training_job_name="my-training-job") + + assert model_trainer.output_data == ( + f"{DEFAULT_OUTPUT_DATA_CONFIG.s3_output_path}" "/my-training-job/output/output.tar.gz" + ) + + +def test_output_data_raises_when_no_training_job(model_trainer): + """Accessing output_data before training raises a clear error.""" + assert model_trainer._latest_training_job is None + with pytest.raises(ValueError, match="No training job"): + _ = model_trainer.output_data + + +def test_output_data_raises_when_no_output_path(model_trainer): + """output_data raises if no S3 output path can be resolved.""" + model_trainer._latest_training_job = TrainingJob(training_job_name="my-training-job") + model_trainer.output_data_config = None + with pytest.raises(ValueError, match="output S3 path"): + _ = model_trainer.output_data + + +def test_output_data_strips_trailing_slash_on_fallback(model_trainer): + """The trailing slash is also normalized when using the trainer fallback.""" + model_trainer._latest_training_job = TrainingJob(training_job_name="my-training-job") + model_trainer.output_data_config = OutputDataConfig( + s3_output_path=f"s3://{DEFAULT_BUCKET}/{DEFAULT_BUCKET_PREFIX}/" + ) + + assert model_trainer.output_data == ( + f"s3://{DEFAULT_BUCKET}/{DEFAULT_BUCKET_PREFIX}" "/my-training-job/output/output.tar.gz" + )