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
51 changes: 51 additions & 0 deletions sagemaker-train/src/sagemaker/train/model_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

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