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
181 changes: 181 additions & 0 deletions sagemaker-core/src/sagemaker/core/training/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -13,8 +13,12 @@
"""Training utilities."""
from __future__ import absolute_import

import io
import json
import os
import tarfile
from typing import Any, Literal
from urllib.parse import urlparse
from sagemaker.core.utils.utils import Unassigned


Expand DownExpand Up@@ -75,3 +79,180 @@ def _is_valid_s3_uri(path: str, path_type: Literal["File", "Directory", "Any"] =
return path.endswith("/")

return path_type == "Any"


_MANIFEST_CHECKPOINT_KEY = "checkpoint_s3_bucket"


def build_nova_hyperpod_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the HyperPod manifest.json S3 URI for a Nova training job.

HyperPod jobs write the manifest directly under the job directory:
``<s3_output_path>/<training_job_name>/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/manifest.json"


def build_nova_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the serverless manifest.json S3 URI for a Nova training job.

Serverless jobs write the manifest under a nested output directory:
``<s3_output_path>/<training_job_name>/output/output/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output/manifest.json"

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The output path for manifest json may be different for serverless and hyperpod.

For serverless: training_job_name/output/output/manifest.json
For SMHP: training_job_name/manifest.json

Can we test this for both flows.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for calling this out! I will push a new revision to address this gap.

I will test all three flows - Serverless, Serverful and HP



def build_nova_output_tar_gz_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the output.tar.gz S3 URI for a Nova training job.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's output.tar.gz.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output.tar.gz"


def _split_s3_uri(s3_uri: str) -> tuple:
"""Split an S3 URI into (bucket, key)."""
parsed = urlparse(s3_uri)
return parsed.netloc, parsed.path.lstrip("/")


def read_nova_checkpoint_uri_from_manifest(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a raw manifest.json object in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the manifest.json object.

Returns:
The ``checkpoint_s3_bucket`` value from the manifest.

Raises:
ValueError: If the object is missing, unparseable, or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
try:
response = s3_client.get_object(Bucket=bucket, Key=key)
manifest = json.loads(response["Body"].read().decode("utf-8"))
except s3_client.exceptions.NoSuchKey:
raise ValueError(f"manifest.json not found at s3://{bucket}/{key}")
except json.JSONDecodeError as e:
raise ValueError(f"Failed to parse manifest.json: {e}")

checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if not checkpoint_uri:
raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json. "
f"Available keys: {list(manifest.keys())}"
)
return checkpoint_uri


def _read_checkpoint_uri_from_tar_gz(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a manifest.json inside an output.tar.gz in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the output.tar.gz object.

Returns:
The ``checkpoint_s3_bucket`` value from the embedded manifest.

Raises:
ValueError: If the archive or manifest is missing or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
response = s3_client.get_object(Bucket=bucket, Key=key)
body = response["Body"].read()
with tarfile.open(fileobj=io.BytesIO(body), mode="r:gz") as tar:
for member in tar.getmembers():
if not member.name.endswith("manifest.json"):
continue
extracted = tar.extractfile(member)
if extracted is None:
continue
manifest = json.loads(extracted.read().decode("utf-8"))
checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if checkpoint_uri:
return checkpoint_uri

raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json within "
f"s3://{bucket}/{key}"
)


def resolve_nova_checkpoint_uri(
s3_client,
s3_output_path: str,
training_job_name: str,
) -> str:
"""Resolve the Nova checkpoint (escrow) URI from a training job's output.

Reads ``checkpoint_s3_bucket`` from the job's manifest.json. The manifest is
first looked up as a raw object, and if that fails, it falls back to the copy
packaged inside ``output.tar.gz``.

Args:
s3_client: A boto3 S3 client.
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
The checkpoint URI recorded in the manifest.

Raises:
ValueError: If the checkpoint URI cannot be resolved from any known
output layout.
"""
# Nova jobs write their manifest to different locations depending on the
# training platform:
# HyperPod: <output>/<job>/manifest.json
# Serverless: <output>/<job>/output/output/manifest.json
# Serverful: <output>/<job>/output/output.tar.gz (manifest is inside)
# Try each in turn and surface every failure if none resolve, so the real
# cause is not masked by a misleading message from the last attempt.
hyperpod_manifest_uri = build_nova_hyperpod_manifest_s3_uri(
s3_output_path, training_job_name
)
serverless_manifest_uri = build_nova_manifest_s3_uri(s3_output_path, training_job_name)
tar_gz_uri = build_nova_output_tar_gz_s3_uri(s3_output_path, training_job_name)

attempts = [
("HyperPod manifest.json", hyperpod_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverless manifest.json", serverless_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverful output.tar.gz", tar_gz_uri, _read_checkpoint_uri_from_tar_gz),
]

errors = []
for label, uri, reader in attempts:
try:
return reader(s3_client, uri)
except Exception as error: # noqa: PERF203 - each attempt may fail independently
errors.append(f"{label} at {uri} failed: {error}")

raise ValueError(
"Could not resolve the Nova checkpoint URI from any known output layout. "
+ " ".join(errors)
)
198 changes: 198 additions & 0 deletions sagemaker-core/tests/unit/test_training_utils.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,198 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Unit tests for Nova manifest/checkpoint helpers in training/utils.py."""
import io
import json
import tarfile

import pytest
from unittest.mock import Mock

from sagemaker.core.training.utils import (
build_nova_hyperpod_manifest_s3_uri,
build_nova_manifest_s3_uri,
build_nova_output_tar_gz_s3_uri,
read_nova_checkpoint_uri_from_manifest,
resolve_nova_checkpoint_uri,
)

CHECKPOINT_URI = "s3://bucket/ckpt/step_100"


def test_build_nova_manifest_s3_uri():
result = build_nova_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_manifest_s3_uri_strips_trailing_slash():
assert build_nova_manifest_s3_uri(
"s3://bucket/output//", "my-job"
) == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_hyperpod_manifest_s3_uri():
result = build_nova_hyperpod_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/manifest.json"


def test_build_nova_output_tar_gz_s3_uri():
result = build_nova_output_tar_gz_s3_uri("s3://bucket/output", "my-job")
assert result == "s3://bucket/output/my-job/output/output.tar.gz"


def _s3_client_returning(body_bytes):
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
body = Mock()
body.read.return_value = body_bytes
client.get_object.return_value = {"Body": body}
return client


def test_read_manifest_returns_checkpoint_uri():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = read_nova_checkpoint_uri_from_manifest(
client, "s3://bucket/output/my-job/output/output/manifest.json"
)
assert result == CHECKPOINT_URI
client.get_object.assert_called_once_with(
Bucket="bucket", Key="output/my-job/output/output/manifest.json"
)


def test_read_manifest_missing_key_raises():
client = _s3_client_returning(json.dumps({"other": "value"}).encode("utf-8"))
with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_not_found_raises():
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
client.get_object.side_effect = client.exceptions.NoSuchKey()
with pytest.raises(ValueError, match="manifest.json not found"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_invalid_json_raises():
client = _s3_client_returning(b"not-json")
with pytest.raises(ValueError, match="Failed to parse manifest.json"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def _make_tar_gz_with_manifest(manifest_dict):
content = json.dumps(manifest_dict).encode("utf-8")
buf = io.BytesIO()
with tarfile.open(fileobj=buf, mode="w:gz") as tar:
info = tarfile.TarInfo(name="manifest.json")
info.size = len(content)
tar.addfile(info, io.BytesIO(content))
return buf.getvalue()


def test_resolve_checkpoint_uri_from_raw_manifest():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_hyperpod_layout():
"""HyperPod writes the manifest at <output>/<job>/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
hyperpod_key = "output/my-job/manifest.json"

def get_object(Bucket, Key):
if Key != hyperpod_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_serverless_layout():
"""Serverless writes the manifest at <output>/<job>/output/output/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
serverless_key = "output/my-job/output/output/manifest.json"

def get_object(Bucket, Key):
if Key != serverless_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_falls_back_to_tar_gz():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"checkpoint_s3_bucket": CHECKPOINT_URI})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_raises_when_both_sources_fail():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"other": "value"})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
Loading
, '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
181 changes: 181 additions & 0 deletions sagemaker-core/src/sagemaker/core/training/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -13,8 +13,12 @@
"""Training utilities."""
from __future__ import absolute_import

import io
import json
import os
import tarfile
from typing import Any, Literal
from urllib.parse import urlparse
from sagemaker.core.utils.utils import Unassigned


Expand DownExpand Up@@ -75,3 +79,180 @@ def _is_valid_s3_uri(path: str, path_type: Literal["File", "Directory", "Any"] =
return path.endswith("/")

return path_type == "Any"


_MANIFEST_CHECKPOINT_KEY = "checkpoint_s3_bucket"


def build_nova_hyperpod_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the HyperPod manifest.json S3 URI for a Nova training job.

HyperPod jobs write the manifest directly under the job directory:
``<s3_output_path>/<training_job_name>/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/manifest.json"


def build_nova_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the serverless manifest.json S3 URI for a Nova training job.

Serverless jobs write the manifest under a nested output directory:
``<s3_output_path>/<training_job_name>/output/output/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output/manifest.json"

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The output path for manifest json may be different for serverless and hyperpod.

For serverless: training_job_name/output/output/manifest.json
For SMHP: training_job_name/manifest.json

Can we test this for both flows.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for calling this out! I will push a new revision to address this gap.

I will test all three flows - Serverless, Serverful and HP



def build_nova_output_tar_gz_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the output.tar.gz S3 URI for a Nova training job.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's output.tar.gz.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output.tar.gz"


def _split_s3_uri(s3_uri: str) -> tuple:
"""Split an S3 URI into (bucket, key)."""
parsed = urlparse(s3_uri)
return parsed.netloc, parsed.path.lstrip("/")


def read_nova_checkpoint_uri_from_manifest(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a raw manifest.json object in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the manifest.json object.

Returns:
The ``checkpoint_s3_bucket`` value from the manifest.

Raises:
ValueError: If the object is missing, unparseable, or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
try:
response = s3_client.get_object(Bucket=bucket, Key=key)
manifest = json.loads(response["Body"].read().decode("utf-8"))
except s3_client.exceptions.NoSuchKey:
raise ValueError(f"manifest.json not found at s3://{bucket}/{key}")
except json.JSONDecodeError as e:
raise ValueError(f"Failed to parse manifest.json: {e}")

checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if not checkpoint_uri:
raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json. "
f"Available keys: {list(manifest.keys())}"
)
return checkpoint_uri


def _read_checkpoint_uri_from_tar_gz(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a manifest.json inside an output.tar.gz in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the output.tar.gz object.

Returns:
The ``checkpoint_s3_bucket`` value from the embedded manifest.

Raises:
ValueError: If the archive or manifest is missing or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
response = s3_client.get_object(Bucket=bucket, Key=key)
body = response["Body"].read()
with tarfile.open(fileobj=io.BytesIO(body), mode="r:gz") as tar:
for member in tar.getmembers():
if not member.name.endswith("manifest.json"):
continue
extracted = tar.extractfile(member)
if extracted is None:
continue
manifest = json.loads(extracted.read().decode("utf-8"))
checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if checkpoint_uri:
return checkpoint_uri

raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json within "
f"s3://{bucket}/{key}"
)


def resolve_nova_checkpoint_uri(
s3_client,
s3_output_path: str,
training_job_name: str,
) -> str:
"""Resolve the Nova checkpoint (escrow) URI from a training job's output.

Reads ``checkpoint_s3_bucket`` from the job's manifest.json. The manifest is
first looked up as a raw object, and if that fails, it falls back to the copy
packaged inside ``output.tar.gz``.

Args:
s3_client: A boto3 S3 client.
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
The checkpoint URI recorded in the manifest.

Raises:
ValueError: If the checkpoint URI cannot be resolved from any known
output layout.
"""
# Nova jobs write their manifest to different locations depending on the
# training platform:
# HyperPod: <output>/<job>/manifest.json
# Serverless: <output>/<job>/output/output/manifest.json
# Serverful: <output>/<job>/output/output.tar.gz (manifest is inside)
# Try each in turn and surface every failure if none resolve, so the real
# cause is not masked by a misleading message from the last attempt.
hyperpod_manifest_uri = build_nova_hyperpod_manifest_s3_uri(
s3_output_path, training_job_name
)
serverless_manifest_uri = build_nova_manifest_s3_uri(s3_output_path, training_job_name)
tar_gz_uri = build_nova_output_tar_gz_s3_uri(s3_output_path, training_job_name)

attempts = [
("HyperPod manifest.json", hyperpod_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverless manifest.json", serverless_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverful output.tar.gz", tar_gz_uri, _read_checkpoint_uri_from_tar_gz),
]

errors = []
for label, uri, reader in attempts:
try:
return reader(s3_client, uri)
except Exception as error: # noqa: PERF203 - each attempt may fail independently
errors.append(f"{label} at {uri} failed: {error}")

raise ValueError(
"Could not resolve the Nova checkpoint URI from any known output layout. "
+ " ".join(errors)
)
198 changes: 198 additions & 0 deletions sagemaker-core/tests/unit/test_training_utils.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,198 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Unit tests for Nova manifest/checkpoint helpers in training/utils.py."""
import io
import json
import tarfile

import pytest
from unittest.mock import Mock

from sagemaker.core.training.utils import (
build_nova_hyperpod_manifest_s3_uri,
build_nova_manifest_s3_uri,
build_nova_output_tar_gz_s3_uri,
read_nova_checkpoint_uri_from_manifest,
resolve_nova_checkpoint_uri,
)

CHECKPOINT_URI = "s3://bucket/ckpt/step_100"


def test_build_nova_manifest_s3_uri():
result = build_nova_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_manifest_s3_uri_strips_trailing_slash():
assert build_nova_manifest_s3_uri(
"s3://bucket/output//", "my-job"
) == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_hyperpod_manifest_s3_uri():
result = build_nova_hyperpod_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/manifest.json"


def test_build_nova_output_tar_gz_s3_uri():
result = build_nova_output_tar_gz_s3_uri("s3://bucket/output", "my-job")
assert result == "s3://bucket/output/my-job/output/output.tar.gz"


def _s3_client_returning(body_bytes):
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
body = Mock()
body.read.return_value = body_bytes
client.get_object.return_value = {"Body": body}
return client


def test_read_manifest_returns_checkpoint_uri():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = read_nova_checkpoint_uri_from_manifest(
client, "s3://bucket/output/my-job/output/output/manifest.json"
)
assert result == CHECKPOINT_URI
client.get_object.assert_called_once_with(
Bucket="bucket", Key="output/my-job/output/output/manifest.json"
)


def test_read_manifest_missing_key_raises():
client = _s3_client_returning(json.dumps({"other": "value"}).encode("utf-8"))
with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_not_found_raises():
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
client.get_object.side_effect = client.exceptions.NoSuchKey()
with pytest.raises(ValueError, match="manifest.json not found"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_invalid_json_raises():
client = _s3_client_returning(b"not-json")
with pytest.raises(ValueError, match="Failed to parse manifest.json"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def _make_tar_gz_with_manifest(manifest_dict):
content = json.dumps(manifest_dict).encode("utf-8")
buf = io.BytesIO()
with tarfile.open(fileobj=buf, mode="w:gz") as tar:
info = tarfile.TarInfo(name="manifest.json")
info.size = len(content)
tar.addfile(info, io.BytesIO(content))
return buf.getvalue()


def test_resolve_checkpoint_uri_from_raw_manifest():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_hyperpod_layout():
"""HyperPod writes the manifest at <output>/<job>/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
hyperpod_key = "output/my-job/manifest.json"

def get_object(Bucket, Key):
if Key != hyperpod_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_serverless_layout():
"""Serverless writes the manifest at <output>/<job>/output/output/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
serverless_key = "output/my-job/output/output/manifest.json"

def get_object(Bucket, Key):
if Key != serverless_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_falls_back_to_tar_gz():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"checkpoint_s3_bucket": CHECKPOINT_URI})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_raises_when_both_sources_fail():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"other": "value"})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
Loading
, '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
181 changes: 181 additions & 0 deletions sagemaker-core/src/sagemaker/core/training/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -13,8 +13,12 @@
"""Training utilities."""
from __future__ import absolute_import

import io
import json
import os
import tarfile
from typing import Any, Literal
from urllib.parse import urlparse
from sagemaker.core.utils.utils import Unassigned


Expand DownExpand Up@@ -75,3 +79,180 @@ def _is_valid_s3_uri(path: str, path_type: Literal["File", "Directory", "Any"] =
return path.endswith("/")

return path_type == "Any"


_MANIFEST_CHECKPOINT_KEY = "checkpoint_s3_bucket"


def build_nova_hyperpod_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the HyperPod manifest.json S3 URI for a Nova training job.

HyperPod jobs write the manifest directly under the job directory:
``<s3_output_path>/<training_job_name>/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/manifest.json"


def build_nova_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the serverless manifest.json S3 URI for a Nova training job.

Serverless jobs write the manifest under a nested output directory:
``<s3_output_path>/<training_job_name>/output/output/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output/manifest.json"

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The output path for manifest json may be different for serverless and hyperpod.

For serverless: training_job_name/output/output/manifest.json
For SMHP: training_job_name/manifest.json

Can we test this for both flows.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for calling this out! I will push a new revision to address this gap.

I will test all three flows - Serverless, Serverful and HP



def build_nova_output_tar_gz_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the output.tar.gz S3 URI for a Nova training job.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's output.tar.gz.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output.tar.gz"


def _split_s3_uri(s3_uri: str) -> tuple:
"""Split an S3 URI into (bucket, key)."""
parsed = urlparse(s3_uri)
return parsed.netloc, parsed.path.lstrip("/")


def read_nova_checkpoint_uri_from_manifest(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a raw manifest.json object in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the manifest.json object.

Returns:
The ``checkpoint_s3_bucket`` value from the manifest.

Raises:
ValueError: If the object is missing, unparseable, or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
try:
response = s3_client.get_object(Bucket=bucket, Key=key)
manifest = json.loads(response["Body"].read().decode("utf-8"))
except s3_client.exceptions.NoSuchKey:
raise ValueError(f"manifest.json not found at s3://{bucket}/{key}")
except json.JSONDecodeError as e:
raise ValueError(f"Failed to parse manifest.json: {e}")

checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if not checkpoint_uri:
raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json. "
f"Available keys: {list(manifest.keys())}"
)
return checkpoint_uri


def _read_checkpoint_uri_from_tar_gz(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a manifest.json inside an output.tar.gz in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the output.tar.gz object.

Returns:
The ``checkpoint_s3_bucket`` value from the embedded manifest.

Raises:
ValueError: If the archive or manifest is missing or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
response = s3_client.get_object(Bucket=bucket, Key=key)
body = response["Body"].read()
with tarfile.open(fileobj=io.BytesIO(body), mode="r:gz") as tar:
for member in tar.getmembers():
if not member.name.endswith("manifest.json"):
continue
extracted = tar.extractfile(member)
if extracted is None:
continue
manifest = json.loads(extracted.read().decode("utf-8"))
checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if checkpoint_uri:
return checkpoint_uri

raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json within "
f"s3://{bucket}/{key}"
)


def resolve_nova_checkpoint_uri(
s3_client,
s3_output_path: str,
training_job_name: str,
) -> str:
"""Resolve the Nova checkpoint (escrow) URI from a training job's output.

Reads ``checkpoint_s3_bucket`` from the job's manifest.json. The manifest is
first looked up as a raw object, and if that fails, it falls back to the copy
packaged inside ``output.tar.gz``.

Args:
s3_client: A boto3 S3 client.
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
The checkpoint URI recorded in the manifest.

Raises:
ValueError: If the checkpoint URI cannot be resolved from any known
output layout.
"""
# Nova jobs write their manifest to different locations depending on the
# training platform:
# HyperPod: <output>/<job>/manifest.json
# Serverless: <output>/<job>/output/output/manifest.json
# Serverful: <output>/<job>/output/output.tar.gz (manifest is inside)
# Try each in turn and surface every failure if none resolve, so the real
# cause is not masked by a misleading message from the last attempt.
hyperpod_manifest_uri = build_nova_hyperpod_manifest_s3_uri(
s3_output_path, training_job_name
)
serverless_manifest_uri = build_nova_manifest_s3_uri(s3_output_path, training_job_name)
tar_gz_uri = build_nova_output_tar_gz_s3_uri(s3_output_path, training_job_name)

attempts = [
("HyperPod manifest.json", hyperpod_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverless manifest.json", serverless_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverful output.tar.gz", tar_gz_uri, _read_checkpoint_uri_from_tar_gz),
]

errors = []
for label, uri, reader in attempts:
try:
return reader(s3_client, uri)
except Exception as error: # noqa: PERF203 - each attempt may fail independently
errors.append(f"{label} at {uri} failed: {error}")

raise ValueError(
"Could not resolve the Nova checkpoint URI from any known output layout. "
+ " ".join(errors)
)
198 changes: 198 additions & 0 deletions sagemaker-core/tests/unit/test_training_utils.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,198 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Unit tests for Nova manifest/checkpoint helpers in training/utils.py."""
import io
import json
import tarfile

import pytest
from unittest.mock import Mock

from sagemaker.core.training.utils import (
build_nova_hyperpod_manifest_s3_uri,
build_nova_manifest_s3_uri,
build_nova_output_tar_gz_s3_uri,
read_nova_checkpoint_uri_from_manifest,
resolve_nova_checkpoint_uri,
)

CHECKPOINT_URI = "s3://bucket/ckpt/step_100"


def test_build_nova_manifest_s3_uri():
result = build_nova_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_manifest_s3_uri_strips_trailing_slash():
assert build_nova_manifest_s3_uri(
"s3://bucket/output//", "my-job"
) == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_hyperpod_manifest_s3_uri():
result = build_nova_hyperpod_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/manifest.json"


def test_build_nova_output_tar_gz_s3_uri():
result = build_nova_output_tar_gz_s3_uri("s3://bucket/output", "my-job")
assert result == "s3://bucket/output/my-job/output/output.tar.gz"


def _s3_client_returning(body_bytes):
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
body = Mock()
body.read.return_value = body_bytes
client.get_object.return_value = {"Body": body}
return client


def test_read_manifest_returns_checkpoint_uri():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = read_nova_checkpoint_uri_from_manifest(
client, "s3://bucket/output/my-job/output/output/manifest.json"
)
assert result == CHECKPOINT_URI
client.get_object.assert_called_once_with(
Bucket="bucket", Key="output/my-job/output/output/manifest.json"
)


def test_read_manifest_missing_key_raises():
client = _s3_client_returning(json.dumps({"other": "value"}).encode("utf-8"))
with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_not_found_raises():
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
client.get_object.side_effect = client.exceptions.NoSuchKey()
with pytest.raises(ValueError, match="manifest.json not found"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_invalid_json_raises():
client = _s3_client_returning(b"not-json")
with pytest.raises(ValueError, match="Failed to parse manifest.json"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def _make_tar_gz_with_manifest(manifest_dict):
content = json.dumps(manifest_dict).encode("utf-8")
buf = io.BytesIO()
with tarfile.open(fileobj=buf, mode="w:gz") as tar:
info = tarfile.TarInfo(name="manifest.json")
info.size = len(content)
tar.addfile(info, io.BytesIO(content))
return buf.getvalue()


def test_resolve_checkpoint_uri_from_raw_manifest():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_hyperpod_layout():
"""HyperPod writes the manifest at <output>/<job>/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
hyperpod_key = "output/my-job/manifest.json"

def get_object(Bucket, Key):
if Key != hyperpod_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_serverless_layout():
"""Serverless writes the manifest at <output>/<job>/output/output/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
serverless_key = "output/my-job/output/output/manifest.json"

def get_object(Bucket, Key):
if Key != serverless_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_falls_back_to_tar_gz():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"checkpoint_s3_bucket": CHECKPOINT_URI})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_raises_when_both_sources_fail():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"other": "value"})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
Loading
, '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
181 changes: 181 additions & 0 deletions sagemaker-core/src/sagemaker/core/training/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -13,8 +13,12 @@
"""Training utilities."""
from __future__ import absolute_import

import io
import json
import os
import tarfile
from typing import Any, Literal
from urllib.parse import urlparse
from sagemaker.core.utils.utils import Unassigned


Expand DownExpand Up@@ -75,3 +79,180 @@ def _is_valid_s3_uri(path: str, path_type: Literal["File", "Directory", "Any"] =
return path.endswith("/")

return path_type == "Any"


_MANIFEST_CHECKPOINT_KEY = "checkpoint_s3_bucket"


def build_nova_hyperpod_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the HyperPod manifest.json S3 URI for a Nova training job.

HyperPod jobs write the manifest directly under the job directory:
``<s3_output_path>/<training_job_name>/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/manifest.json"


def build_nova_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the serverless manifest.json S3 URI for a Nova training job.

Serverless jobs write the manifest under a nested output directory:
``<s3_output_path>/<training_job_name>/output/output/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output/manifest.json"

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The output path for manifest json may be different for serverless and hyperpod.

For serverless: training_job_name/output/output/manifest.json
For SMHP: training_job_name/manifest.json

Can we test this for both flows.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for calling this out! I will push a new revision to address this gap.

I will test all three flows - Serverless, Serverful and HP



def build_nova_output_tar_gz_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the output.tar.gz S3 URI for a Nova training job.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's output.tar.gz.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output.tar.gz"


def _split_s3_uri(s3_uri: str) -> tuple:
"""Split an S3 URI into (bucket, key)."""
parsed = urlparse(s3_uri)
return parsed.netloc, parsed.path.lstrip("/")


def read_nova_checkpoint_uri_from_manifest(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a raw manifest.json object in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the manifest.json object.

Returns:
The ``checkpoint_s3_bucket`` value from the manifest.

Raises:
ValueError: If the object is missing, unparseable, or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
try:
response = s3_client.get_object(Bucket=bucket, Key=key)
manifest = json.loads(response["Body"].read().decode("utf-8"))
except s3_client.exceptions.NoSuchKey:
raise ValueError(f"manifest.json not found at s3://{bucket}/{key}")
except json.JSONDecodeError as e:
raise ValueError(f"Failed to parse manifest.json: {e}")

checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if not checkpoint_uri:
raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json. "
f"Available keys: {list(manifest.keys())}"
)
return checkpoint_uri


def _read_checkpoint_uri_from_tar_gz(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a manifest.json inside an output.tar.gz in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the output.tar.gz object.

Returns:
The ``checkpoint_s3_bucket`` value from the embedded manifest.

Raises:
ValueError: If the archive or manifest is missing or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
response = s3_client.get_object(Bucket=bucket, Key=key)
body = response["Body"].read()
with tarfile.open(fileobj=io.BytesIO(body), mode="r:gz") as tar:
for member in tar.getmembers():
if not member.name.endswith("manifest.json"):
continue
extracted = tar.extractfile(member)
if extracted is None:
continue
manifest = json.loads(extracted.read().decode("utf-8"))
checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if checkpoint_uri:
return checkpoint_uri

raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json within "
f"s3://{bucket}/{key}"
)


def resolve_nova_checkpoint_uri(
s3_client,
s3_output_path: str,
training_job_name: str,
) -> str:
"""Resolve the Nova checkpoint (escrow) URI from a training job's output.

Reads ``checkpoint_s3_bucket`` from the job's manifest.json. The manifest is
first looked up as a raw object, and if that fails, it falls back to the copy
packaged inside ``output.tar.gz``.

Args:
s3_client: A boto3 S3 client.
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
The checkpoint URI recorded in the manifest.

Raises:
ValueError: If the checkpoint URI cannot be resolved from any known
output layout.
"""
# Nova jobs write their manifest to different locations depending on the
# training platform:
# HyperPod: <output>/<job>/manifest.json
# Serverless: <output>/<job>/output/output/manifest.json
# Serverful: <output>/<job>/output/output.tar.gz (manifest is inside)
# Try each in turn and surface every failure if none resolve, so the real
# cause is not masked by a misleading message from the last attempt.
hyperpod_manifest_uri = build_nova_hyperpod_manifest_s3_uri(
s3_output_path, training_job_name
)
serverless_manifest_uri = build_nova_manifest_s3_uri(s3_output_path, training_job_name)
tar_gz_uri = build_nova_output_tar_gz_s3_uri(s3_output_path, training_job_name)

attempts = [
("HyperPod manifest.json", hyperpod_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverless manifest.json", serverless_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverful output.tar.gz", tar_gz_uri, _read_checkpoint_uri_from_tar_gz),
]

errors = []
for label, uri, reader in attempts:
try:
return reader(s3_client, uri)
except Exception as error: # noqa: PERF203 - each attempt may fail independently
errors.append(f"{label} at {uri} failed: {error}")

raise ValueError(
"Could not resolve the Nova checkpoint URI from any known output layout. "
+ " ".join(errors)
)
198 changes: 198 additions & 0 deletions sagemaker-core/tests/unit/test_training_utils.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,198 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Unit tests for Nova manifest/checkpoint helpers in training/utils.py."""
import io
import json
import tarfile

import pytest
from unittest.mock import Mock

from sagemaker.core.training.utils import (
build_nova_hyperpod_manifest_s3_uri,
build_nova_manifest_s3_uri,
build_nova_output_tar_gz_s3_uri,
read_nova_checkpoint_uri_from_manifest,
resolve_nova_checkpoint_uri,
)

CHECKPOINT_URI = "s3://bucket/ckpt/step_100"


def test_build_nova_manifest_s3_uri():
result = build_nova_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_manifest_s3_uri_strips_trailing_slash():
assert build_nova_manifest_s3_uri(
"s3://bucket/output//", "my-job"
) == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_hyperpod_manifest_s3_uri():
result = build_nova_hyperpod_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/manifest.json"


def test_build_nova_output_tar_gz_s3_uri():
result = build_nova_output_tar_gz_s3_uri("s3://bucket/output", "my-job")
assert result == "s3://bucket/output/my-job/output/output.tar.gz"


def _s3_client_returning(body_bytes):
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
body = Mock()
body.read.return_value = body_bytes
client.get_object.return_value = {"Body": body}
return client


def test_read_manifest_returns_checkpoint_uri():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = read_nova_checkpoint_uri_from_manifest(
client, "s3://bucket/output/my-job/output/output/manifest.json"
)
assert result == CHECKPOINT_URI
client.get_object.assert_called_once_with(
Bucket="bucket", Key="output/my-job/output/output/manifest.json"
)


def test_read_manifest_missing_key_raises():
client = _s3_client_returning(json.dumps({"other": "value"}).encode("utf-8"))
with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_not_found_raises():
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
client.get_object.side_effect = client.exceptions.NoSuchKey()
with pytest.raises(ValueError, match="manifest.json not found"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_invalid_json_raises():
client = _s3_client_returning(b"not-json")
with pytest.raises(ValueError, match="Failed to parse manifest.json"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def _make_tar_gz_with_manifest(manifest_dict):
content = json.dumps(manifest_dict).encode("utf-8")
buf = io.BytesIO()
with tarfile.open(fileobj=buf, mode="w:gz") as tar:
info = tarfile.TarInfo(name="manifest.json")
info.size = len(content)
tar.addfile(info, io.BytesIO(content))
return buf.getvalue()


def test_resolve_checkpoint_uri_from_raw_manifest():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_hyperpod_layout():
"""HyperPod writes the manifest at <output>/<job>/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
hyperpod_key = "output/my-job/manifest.json"

def get_object(Bucket, Key):
if Key != hyperpod_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_serverless_layout():
"""Serverless writes the manifest at <output>/<job>/output/output/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
serverless_key = "output/my-job/output/output/manifest.json"

def get_object(Bucket, Key):
if Key != serverless_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_falls_back_to_tar_gz():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"checkpoint_s3_bucket": CHECKPOINT_URI})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_raises_when_both_sources_fail():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"other": "value"})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
Loading
, '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
181 changes: 181 additions & 0 deletions sagemaker-core/src/sagemaker/core/training/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -13,8 +13,12 @@
"""Training utilities."""
from __future__ import absolute_import

import io
import json
import os
import tarfile
from typing import Any, Literal
from urllib.parse import urlparse
from sagemaker.core.utils.utils import Unassigned


Expand DownExpand Up@@ -75,3 +79,180 @@ def _is_valid_s3_uri(path: str, path_type: Literal["File", "Directory", "Any"] =
return path.endswith("/")

return path_type == "Any"


_MANIFEST_CHECKPOINT_KEY = "checkpoint_s3_bucket"


def build_nova_hyperpod_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the HyperPod manifest.json S3 URI for a Nova training job.

HyperPod jobs write the manifest directly under the job directory:
``<s3_output_path>/<training_job_name>/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/manifest.json"


def build_nova_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the serverless manifest.json S3 URI for a Nova training job.

Serverless jobs write the manifest under a nested output directory:
``<s3_output_path>/<training_job_name>/output/output/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output/manifest.json"

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The output path for manifest json may be different for serverless and hyperpod.

For serverless: training_job_name/output/output/manifest.json
For SMHP: training_job_name/manifest.json

Can we test this for both flows.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for calling this out! I will push a new revision to address this gap.

I will test all three flows - Serverless, Serverful and HP



def build_nova_output_tar_gz_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the output.tar.gz S3 URI for a Nova training job.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's output.tar.gz.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output.tar.gz"


def _split_s3_uri(s3_uri: str) -> tuple:
"""Split an S3 URI into (bucket, key)."""
parsed = urlparse(s3_uri)
return parsed.netloc, parsed.path.lstrip("/")


def read_nova_checkpoint_uri_from_manifest(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a raw manifest.json object in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the manifest.json object.

Returns:
The ``checkpoint_s3_bucket`` value from the manifest.

Raises:
ValueError: If the object is missing, unparseable, or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
try:
response = s3_client.get_object(Bucket=bucket, Key=key)
manifest = json.loads(response["Body"].read().decode("utf-8"))
except s3_client.exceptions.NoSuchKey:
raise ValueError(f"manifest.json not found at s3://{bucket}/{key}")
except json.JSONDecodeError as e:
raise ValueError(f"Failed to parse manifest.json: {e}")

checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if not checkpoint_uri:
raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json. "
f"Available keys: {list(manifest.keys())}"
)
return checkpoint_uri


def _read_checkpoint_uri_from_tar_gz(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a manifest.json inside an output.tar.gz in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the output.tar.gz object.

Returns:
The ``checkpoint_s3_bucket`` value from the embedded manifest.

Raises:
ValueError: If the archive or manifest is missing or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
response = s3_client.get_object(Bucket=bucket, Key=key)
body = response["Body"].read()
with tarfile.open(fileobj=io.BytesIO(body), mode="r:gz") as tar:
for member in tar.getmembers():
if not member.name.endswith("manifest.json"):
continue
extracted = tar.extractfile(member)
if extracted is None:
continue
manifest = json.loads(extracted.read().decode("utf-8"))
checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if checkpoint_uri:
return checkpoint_uri

raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json within "
f"s3://{bucket}/{key}"
)


def resolve_nova_checkpoint_uri(
s3_client,
s3_output_path: str,
training_job_name: str,
) -> str:
"""Resolve the Nova checkpoint (escrow) URI from a training job's output.

Reads ``checkpoint_s3_bucket`` from the job's manifest.json. The manifest is
first looked up as a raw object, and if that fails, it falls back to the copy
packaged inside ``output.tar.gz``.

Args:
s3_client: A boto3 S3 client.
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
The checkpoint URI recorded in the manifest.

Raises:
ValueError: If the checkpoint URI cannot be resolved from any known
output layout.
"""
# Nova jobs write their manifest to different locations depending on the
# training platform:
# HyperPod: <output>/<job>/manifest.json
# Serverless: <output>/<job>/output/output/manifest.json
# Serverful: <output>/<job>/output/output.tar.gz (manifest is inside)
# Try each in turn and surface every failure if none resolve, so the real
# cause is not masked by a misleading message from the last attempt.
hyperpod_manifest_uri = build_nova_hyperpod_manifest_s3_uri(
s3_output_path, training_job_name
)
serverless_manifest_uri = build_nova_manifest_s3_uri(s3_output_path, training_job_name)
tar_gz_uri = build_nova_output_tar_gz_s3_uri(s3_output_path, training_job_name)

attempts = [
("HyperPod manifest.json", hyperpod_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverless manifest.json", serverless_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverful output.tar.gz", tar_gz_uri, _read_checkpoint_uri_from_tar_gz),
]

errors = []
for label, uri, reader in attempts:
try:
return reader(s3_client, uri)
except Exception as error: # noqa: PERF203 - each attempt may fail independently
errors.append(f"{label} at {uri} failed: {error}")

raise ValueError(
"Could not resolve the Nova checkpoint URI from any known output layout. "
+ " ".join(errors)
)
198 changes: 198 additions & 0 deletions sagemaker-core/tests/unit/test_training_utils.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,198 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Unit tests for Nova manifest/checkpoint helpers in training/utils.py."""
import io
import json
import tarfile

import pytest
from unittest.mock import Mock

from sagemaker.core.training.utils import (
build_nova_hyperpod_manifest_s3_uri,
build_nova_manifest_s3_uri,
build_nova_output_tar_gz_s3_uri,
read_nova_checkpoint_uri_from_manifest,
resolve_nova_checkpoint_uri,
)

CHECKPOINT_URI = "s3://bucket/ckpt/step_100"


def test_build_nova_manifest_s3_uri():
result = build_nova_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_manifest_s3_uri_strips_trailing_slash():
assert build_nova_manifest_s3_uri(
"s3://bucket/output//", "my-job"
) == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_hyperpod_manifest_s3_uri():
result = build_nova_hyperpod_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/manifest.json"


def test_build_nova_output_tar_gz_s3_uri():
result = build_nova_output_tar_gz_s3_uri("s3://bucket/output", "my-job")
assert result == "s3://bucket/output/my-job/output/output.tar.gz"


def _s3_client_returning(body_bytes):
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
body = Mock()
body.read.return_value = body_bytes
client.get_object.return_value = {"Body": body}
return client


def test_read_manifest_returns_checkpoint_uri():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = read_nova_checkpoint_uri_from_manifest(
client, "s3://bucket/output/my-job/output/output/manifest.json"
)
assert result == CHECKPOINT_URI
client.get_object.assert_called_once_with(
Bucket="bucket", Key="output/my-job/output/output/manifest.json"
)


def test_read_manifest_missing_key_raises():
client = _s3_client_returning(json.dumps({"other": "value"}).encode("utf-8"))
with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_not_found_raises():
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
client.get_object.side_effect = client.exceptions.NoSuchKey()
with pytest.raises(ValueError, match="manifest.json not found"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_invalid_json_raises():
client = _s3_client_returning(b"not-json")
with pytest.raises(ValueError, match="Failed to parse manifest.json"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def _make_tar_gz_with_manifest(manifest_dict):
content = json.dumps(manifest_dict).encode("utf-8")
buf = io.BytesIO()
with tarfile.open(fileobj=buf, mode="w:gz") as tar:
info = tarfile.TarInfo(name="manifest.json")
info.size = len(content)
tar.addfile(info, io.BytesIO(content))
return buf.getvalue()


def test_resolve_checkpoint_uri_from_raw_manifest():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_hyperpod_layout():
"""HyperPod writes the manifest at <output>/<job>/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
hyperpod_key = "output/my-job/manifest.json"

def get_object(Bucket, Key):
if Key != hyperpod_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_serverless_layout():
"""Serverless writes the manifest at <output>/<job>/output/output/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
serverless_key = "output/my-job/output/output/manifest.json"

def get_object(Bucket, Key):
if Key != serverless_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_falls_back_to_tar_gz():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"checkpoint_s3_bucket": CHECKPOINT_URI})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_raises_when_both_sources_fail():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"other": "value"})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
Loading
, '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
181 changes: 181 additions & 0 deletions sagemaker-core/src/sagemaker/core/training/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -13,8 +13,12 @@
"""Training utilities."""
from __future__ import absolute_import

import io
import json
import os
import tarfile
from typing import Any, Literal
from urllib.parse import urlparse
from sagemaker.core.utils.utils import Unassigned


Expand DownExpand Up@@ -75,3 +79,180 @@ def _is_valid_s3_uri(path: str, path_type: Literal["File", "Directory", "Any"] =
return path.endswith("/")

return path_type == "Any"


_MANIFEST_CHECKPOINT_KEY = "checkpoint_s3_bucket"


def build_nova_hyperpod_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the HyperPod manifest.json S3 URI for a Nova training job.

HyperPod jobs write the manifest directly under the job directory:
``<s3_output_path>/<training_job_name>/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/manifest.json"


def build_nova_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the serverless manifest.json S3 URI for a Nova training job.

Serverless jobs write the manifest under a nested output directory:
``<s3_output_path>/<training_job_name>/output/output/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output/manifest.json"

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The output path for manifest json may be different for serverless and hyperpod.

For serverless: training_job_name/output/output/manifest.json
For SMHP: training_job_name/manifest.json

Can we test this for both flows.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for calling this out! I will push a new revision to address this gap.

I will test all three flows - Serverless, Serverful and HP



def build_nova_output_tar_gz_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the output.tar.gz S3 URI for a Nova training job.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's output.tar.gz.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output.tar.gz"


def _split_s3_uri(s3_uri: str) -> tuple:
"""Split an S3 URI into (bucket, key)."""
parsed = urlparse(s3_uri)
return parsed.netloc, parsed.path.lstrip("/")


def read_nova_checkpoint_uri_from_manifest(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a raw manifest.json object in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the manifest.json object.

Returns:
The ``checkpoint_s3_bucket`` value from the manifest.

Raises:
ValueError: If the object is missing, unparseable, or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
try:
response = s3_client.get_object(Bucket=bucket, Key=key)
manifest = json.loads(response["Body"].read().decode("utf-8"))
except s3_client.exceptions.NoSuchKey:
raise ValueError(f"manifest.json not found at s3://{bucket}/{key}")
except json.JSONDecodeError as e:
raise ValueError(f"Failed to parse manifest.json: {e}")

checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if not checkpoint_uri:
raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json. "
f"Available keys: {list(manifest.keys())}"
)
return checkpoint_uri


def _read_checkpoint_uri_from_tar_gz(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a manifest.json inside an output.tar.gz in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the output.tar.gz object.

Returns:
The ``checkpoint_s3_bucket`` value from the embedded manifest.

Raises:
ValueError: If the archive or manifest is missing or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
response = s3_client.get_object(Bucket=bucket, Key=key)
body = response["Body"].read()
with tarfile.open(fileobj=io.BytesIO(body), mode="r:gz") as tar:
for member in tar.getmembers():
if not member.name.endswith("manifest.json"):
continue
extracted = tar.extractfile(member)
if extracted is None:
continue
manifest = json.loads(extracted.read().decode("utf-8"))
checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if checkpoint_uri:
return checkpoint_uri

raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json within "
f"s3://{bucket}/{key}"
)


def resolve_nova_checkpoint_uri(
s3_client,
s3_output_path: str,
training_job_name: str,
) -> str:
"""Resolve the Nova checkpoint (escrow) URI from a training job's output.

Reads ``checkpoint_s3_bucket`` from the job's manifest.json. The manifest is
first looked up as a raw object, and if that fails, it falls back to the copy
packaged inside ``output.tar.gz``.

Args:
s3_client: A boto3 S3 client.
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
The checkpoint URI recorded in the manifest.

Raises:
ValueError: If the checkpoint URI cannot be resolved from any known
output layout.
"""
# Nova jobs write their manifest to different locations depending on the
# training platform:
# HyperPod: <output>/<job>/manifest.json
# Serverless: <output>/<job>/output/output/manifest.json
# Serverful: <output>/<job>/output/output.tar.gz (manifest is inside)
# Try each in turn and surface every failure if none resolve, so the real
# cause is not masked by a misleading message from the last attempt.
hyperpod_manifest_uri = build_nova_hyperpod_manifest_s3_uri(
s3_output_path, training_job_name
)
serverless_manifest_uri = build_nova_manifest_s3_uri(s3_output_path, training_job_name)
tar_gz_uri = build_nova_output_tar_gz_s3_uri(s3_output_path, training_job_name)

attempts = [
("HyperPod manifest.json", hyperpod_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverless manifest.json", serverless_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverful output.tar.gz", tar_gz_uri, _read_checkpoint_uri_from_tar_gz),
]

errors = []
for label, uri, reader in attempts:
try:
return reader(s3_client, uri)
except Exception as error: # noqa: PERF203 - each attempt may fail independently
errors.append(f"{label} at {uri} failed: {error}")

raise ValueError(
"Could not resolve the Nova checkpoint URI from any known output layout. "
+ " ".join(errors)
)
198 changes: 198 additions & 0 deletions sagemaker-core/tests/unit/test_training_utils.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,198 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Unit tests for Nova manifest/checkpoint helpers in training/utils.py."""
import io
import json
import tarfile

import pytest
from unittest.mock import Mock

from sagemaker.core.training.utils import (
build_nova_hyperpod_manifest_s3_uri,
build_nova_manifest_s3_uri,
build_nova_output_tar_gz_s3_uri,
read_nova_checkpoint_uri_from_manifest,
resolve_nova_checkpoint_uri,
)

CHECKPOINT_URI = "s3://bucket/ckpt/step_100"


def test_build_nova_manifest_s3_uri():
result = build_nova_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_manifest_s3_uri_strips_trailing_slash():
assert build_nova_manifest_s3_uri(
"s3://bucket/output//", "my-job"
) == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_hyperpod_manifest_s3_uri():
result = build_nova_hyperpod_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/manifest.json"


def test_build_nova_output_tar_gz_s3_uri():
result = build_nova_output_tar_gz_s3_uri("s3://bucket/output", "my-job")
assert result == "s3://bucket/output/my-job/output/output.tar.gz"


def _s3_client_returning(body_bytes):
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
body = Mock()
body.read.return_value = body_bytes
client.get_object.return_value = {"Body": body}
return client


def test_read_manifest_returns_checkpoint_uri():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = read_nova_checkpoint_uri_from_manifest(
client, "s3://bucket/output/my-job/output/output/manifest.json"
)
assert result == CHECKPOINT_URI
client.get_object.assert_called_once_with(
Bucket="bucket", Key="output/my-job/output/output/manifest.json"
)


def test_read_manifest_missing_key_raises():
client = _s3_client_returning(json.dumps({"other": "value"}).encode("utf-8"))
with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_not_found_raises():
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
client.get_object.side_effect = client.exceptions.NoSuchKey()
with pytest.raises(ValueError, match="manifest.json not found"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_invalid_json_raises():
client = _s3_client_returning(b"not-json")
with pytest.raises(ValueError, match="Failed to parse manifest.json"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def _make_tar_gz_with_manifest(manifest_dict):
content = json.dumps(manifest_dict).encode("utf-8")
buf = io.BytesIO()
with tarfile.open(fileobj=buf, mode="w:gz") as tar:
info = tarfile.TarInfo(name="manifest.json")
info.size = len(content)
tar.addfile(info, io.BytesIO(content))
return buf.getvalue()


def test_resolve_checkpoint_uri_from_raw_manifest():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_hyperpod_layout():
"""HyperPod writes the manifest at <output>/<job>/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
hyperpod_key = "output/my-job/manifest.json"

def get_object(Bucket, Key):
if Key != hyperpod_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_serverless_layout():
"""Serverless writes the manifest at <output>/<job>/output/output/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
serverless_key = "output/my-job/output/output/manifest.json"

def get_object(Bucket, Key):
if Key != serverless_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_falls_back_to_tar_gz():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"checkpoint_s3_bucket": CHECKPOINT_URI})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_raises_when_both_sources_fail():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"other": "value"})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
Loading
, '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
181 changes: 181 additions & 0 deletions sagemaker-core/src/sagemaker/core/training/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -13,8 +13,12 @@
"""Training utilities."""
from __future__ import absolute_import

import io
import json
import os
import tarfile
from typing import Any, Literal
from urllib.parse import urlparse
from sagemaker.core.utils.utils import Unassigned


Expand DownExpand Up@@ -75,3 +79,180 @@ def _is_valid_s3_uri(path: str, path_type: Literal["File", "Directory", "Any"] =
return path.endswith("/")

return path_type == "Any"


_MANIFEST_CHECKPOINT_KEY = "checkpoint_s3_bucket"


def build_nova_hyperpod_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the HyperPod manifest.json S3 URI for a Nova training job.

HyperPod jobs write the manifest directly under the job directory:
``<s3_output_path>/<training_job_name>/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/manifest.json"


def build_nova_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the serverless manifest.json S3 URI for a Nova training job.

Serverless jobs write the manifest under a nested output directory:
``<s3_output_path>/<training_job_name>/output/output/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output/manifest.json"

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The output path for manifest json may be different for serverless and hyperpod.

For serverless: training_job_name/output/output/manifest.json
For SMHP: training_job_name/manifest.json

Can we test this for both flows.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for calling this out! I will push a new revision to address this gap.

I will test all three flows - Serverless, Serverful and HP



def build_nova_output_tar_gz_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the output.tar.gz S3 URI for a Nova training job.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's output.tar.gz.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output.tar.gz"


def _split_s3_uri(s3_uri: str) -> tuple:
"""Split an S3 URI into (bucket, key)."""
parsed = urlparse(s3_uri)
return parsed.netloc, parsed.path.lstrip("/")


def read_nova_checkpoint_uri_from_manifest(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a raw manifest.json object in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the manifest.json object.

Returns:
The ``checkpoint_s3_bucket`` value from the manifest.

Raises:
ValueError: If the object is missing, unparseable, or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
try:
response = s3_client.get_object(Bucket=bucket, Key=key)
manifest = json.loads(response["Body"].read().decode("utf-8"))
except s3_client.exceptions.NoSuchKey:
raise ValueError(f"manifest.json not found at s3://{bucket}/{key}")
except json.JSONDecodeError as e:
raise ValueError(f"Failed to parse manifest.json: {e}")

checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if not checkpoint_uri:
raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json. "
f"Available keys: {list(manifest.keys())}"
)
return checkpoint_uri


def _read_checkpoint_uri_from_tar_gz(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a manifest.json inside an output.tar.gz in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the output.tar.gz object.

Returns:
The ``checkpoint_s3_bucket`` value from the embedded manifest.

Raises:
ValueError: If the archive or manifest is missing or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
response = s3_client.get_object(Bucket=bucket, Key=key)
body = response["Body"].read()
with tarfile.open(fileobj=io.BytesIO(body), mode="r:gz") as tar:
for member in tar.getmembers():
if not member.name.endswith("manifest.json"):
continue
extracted = tar.extractfile(member)
if extracted is None:
continue
manifest = json.loads(extracted.read().decode("utf-8"))
checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if checkpoint_uri:
return checkpoint_uri

raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json within "
f"s3://{bucket}/{key}"
)


def resolve_nova_checkpoint_uri(
s3_client,
s3_output_path: str,
training_job_name: str,
) -> str:
"""Resolve the Nova checkpoint (escrow) URI from a training job's output.

Reads ``checkpoint_s3_bucket`` from the job's manifest.json. The manifest is
first looked up as a raw object, and if that fails, it falls back to the copy
packaged inside ``output.tar.gz``.

Args:
s3_client: A boto3 S3 client.
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
The checkpoint URI recorded in the manifest.

Raises:
ValueError: If the checkpoint URI cannot be resolved from any known
output layout.
"""
# Nova jobs write their manifest to different locations depending on the
# training platform:
# HyperPod: <output>/<job>/manifest.json
# Serverless: <output>/<job>/output/output/manifest.json
# Serverful: <output>/<job>/output/output.tar.gz (manifest is inside)
# Try each in turn and surface every failure if none resolve, so the real
# cause is not masked by a misleading message from the last attempt.
hyperpod_manifest_uri = build_nova_hyperpod_manifest_s3_uri(
s3_output_path, training_job_name
)
serverless_manifest_uri = build_nova_manifest_s3_uri(s3_output_path, training_job_name)
tar_gz_uri = build_nova_output_tar_gz_s3_uri(s3_output_path, training_job_name)

attempts = [
("HyperPod manifest.json", hyperpod_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverless manifest.json", serverless_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverful output.tar.gz", tar_gz_uri, _read_checkpoint_uri_from_tar_gz),
]

errors = []
for label, uri, reader in attempts:
try:
return reader(s3_client, uri)
except Exception as error: # noqa: PERF203 - each attempt may fail independently
errors.append(f"{label} at {uri} failed: {error}")

raise ValueError(
"Could not resolve the Nova checkpoint URI from any known output layout. "
+ " ".join(errors)
)
198 changes: 198 additions & 0 deletions sagemaker-core/tests/unit/test_training_utils.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,198 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Unit tests for Nova manifest/checkpoint helpers in training/utils.py."""
import io
import json
import tarfile

import pytest
from unittest.mock import Mock

from sagemaker.core.training.utils import (
build_nova_hyperpod_manifest_s3_uri,
build_nova_manifest_s3_uri,
build_nova_output_tar_gz_s3_uri,
read_nova_checkpoint_uri_from_manifest,
resolve_nova_checkpoint_uri,
)

CHECKPOINT_URI = "s3://bucket/ckpt/step_100"


def test_build_nova_manifest_s3_uri():
result = build_nova_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_manifest_s3_uri_strips_trailing_slash():
assert build_nova_manifest_s3_uri(
"s3://bucket/output//", "my-job"
) == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_hyperpod_manifest_s3_uri():
result = build_nova_hyperpod_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/manifest.json"


def test_build_nova_output_tar_gz_s3_uri():
result = build_nova_output_tar_gz_s3_uri("s3://bucket/output", "my-job")
assert result == "s3://bucket/output/my-job/output/output.tar.gz"


def _s3_client_returning(body_bytes):
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
body = Mock()
body.read.return_value = body_bytes
client.get_object.return_value = {"Body": body}
return client


def test_read_manifest_returns_checkpoint_uri():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = read_nova_checkpoint_uri_from_manifest(
client, "s3://bucket/output/my-job/output/output/manifest.json"
)
assert result == CHECKPOINT_URI
client.get_object.assert_called_once_with(
Bucket="bucket", Key="output/my-job/output/output/manifest.json"
)


def test_read_manifest_missing_key_raises():
client = _s3_client_returning(json.dumps({"other": "value"}).encode("utf-8"))
with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_not_found_raises():
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
client.get_object.side_effect = client.exceptions.NoSuchKey()
with pytest.raises(ValueError, match="manifest.json not found"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_invalid_json_raises():
client = _s3_client_returning(b"not-json")
with pytest.raises(ValueError, match="Failed to parse manifest.json"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def _make_tar_gz_with_manifest(manifest_dict):
content = json.dumps(manifest_dict).encode("utf-8")
buf = io.BytesIO()
with tarfile.open(fileobj=buf, mode="w:gz") as tar:
info = tarfile.TarInfo(name="manifest.json")
info.size = len(content)
tar.addfile(info, io.BytesIO(content))
return buf.getvalue()


def test_resolve_checkpoint_uri_from_raw_manifest():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_hyperpod_layout():
"""HyperPod writes the manifest at <output>/<job>/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
hyperpod_key = "output/my-job/manifest.json"

def get_object(Bucket, Key):
if Key != hyperpod_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_serverless_layout():
"""Serverless writes the manifest at <output>/<job>/output/output/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
serverless_key = "output/my-job/output/output/manifest.json"

def get_object(Bucket, Key):
if Key != serverless_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_falls_back_to_tar_gz():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"checkpoint_s3_bucket": CHECKPOINT_URI})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_raises_when_both_sources_fail():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"other": "value"})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
Loading
, '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
181 changes: 181 additions & 0 deletions sagemaker-core/src/sagemaker/core/training/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -13,8 +13,12 @@
"""Training utilities."""
from __future__ import absolute_import

import io
import json
import os
import tarfile
from typing import Any, Literal
from urllib.parse import urlparse
from sagemaker.core.utils.utils import Unassigned


Expand DownExpand Up@@ -75,3 +79,180 @@ def _is_valid_s3_uri(path: str, path_type: Literal["File", "Directory", "Any"] =
return path.endswith("/")

return path_type == "Any"


_MANIFEST_CHECKPOINT_KEY = "checkpoint_s3_bucket"


def build_nova_hyperpod_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the HyperPod manifest.json S3 URI for a Nova training job.

HyperPod jobs write the manifest directly under the job directory:
``<s3_output_path>/<training_job_name>/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/manifest.json"


def build_nova_manifest_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the serverless manifest.json S3 URI for a Nova training job.

Serverless jobs write the manifest under a nested output directory:
``<s3_output_path>/<training_job_name>/output/output/manifest.json``.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's manifest.json.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output/manifest.json"

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The output path for manifest json may be different for serverless and hyperpod.

For serverless: training_job_name/output/output/manifest.json
For SMHP: training_job_name/manifest.json

Can we test this for both flows.

Copy link
Copy Markdown
ContributorAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for calling this out! I will push a new revision to address this gap.

I will test all three flows - Serverless, Serverful and HP



def build_nova_output_tar_gz_s3_uri(s3_output_path: str, training_job_name: str) -> str:
"""Build the output.tar.gz S3 URI for a Nova training job.

Args:
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
Fully-qualified S3 URI to the job's output.tar.gz.
"""
output_path = s3_output_path.rstrip("/")
return f"{output_path}/{training_job_name}/output/output.tar.gz"


def _split_s3_uri(s3_uri: str) -> tuple:
"""Split an S3 URI into (bucket, key)."""
parsed = urlparse(s3_uri)
return parsed.netloc, parsed.path.lstrip("/")


def read_nova_checkpoint_uri_from_manifest(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a raw manifest.json object in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the manifest.json object.

Returns:
The ``checkpoint_s3_bucket`` value from the manifest.

Raises:
ValueError: If the object is missing, unparseable, or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
try:
response = s3_client.get_object(Bucket=bucket, Key=key)
manifest = json.loads(response["Body"].read().decode("utf-8"))
except s3_client.exceptions.NoSuchKey:
raise ValueError(f"manifest.json not found at s3://{bucket}/{key}")
except json.JSONDecodeError as e:
raise ValueError(f"Failed to parse manifest.json: {e}")

checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if not checkpoint_uri:
raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json. "
f"Available keys: {list(manifest.keys())}"
)
return checkpoint_uri


def _read_checkpoint_uri_from_tar_gz(s3_client, s3_uri: str) -> str:
"""Read the checkpoint URI from a manifest.json inside an output.tar.gz in S3.

Args:
s3_client: A boto3 S3 client.
s3_uri: S3 URI of the output.tar.gz object.

Returns:
The ``checkpoint_s3_bucket`` value from the embedded manifest.

Raises:
ValueError: If the archive or manifest is missing or lacks the key.
"""
bucket, key = _split_s3_uri(s3_uri)
response = s3_client.get_object(Bucket=bucket, Key=key)
body = response["Body"].read()
with tarfile.open(fileobj=io.BytesIO(body), mode="r:gz") as tar:
for member in tar.getmembers():
if not member.name.endswith("manifest.json"):
continue
extracted = tar.extractfile(member)
if extracted is None:
continue
manifest = json.loads(extracted.read().decode("utf-8"))
checkpoint_uri = manifest.get(_MANIFEST_CHECKPOINT_KEY)
if checkpoint_uri:
return checkpoint_uri

raise ValueError(
f"'{_MANIFEST_CHECKPOINT_KEY}' not found in manifest.json within "
f"s3://{bucket}/{key}"
)


def resolve_nova_checkpoint_uri(
s3_client,
s3_output_path: str,
training_job_name: str,
) -> str:
"""Resolve the Nova checkpoint (escrow) URI from a training job's output.

Reads ``checkpoint_s3_bucket`` from the job's manifest.json. The manifest is
first looked up as a raw object, and if that fails, it falls back to the copy
packaged inside ``output.tar.gz``.

Args:
s3_client: A boto3 S3 client.
s3_output_path: The training job's ``output_data_config.s3_output_path``.
training_job_name: The training job name.

Returns:
The checkpoint URI recorded in the manifest.

Raises:
ValueError: If the checkpoint URI cannot be resolved from any known
output layout.
"""
# Nova jobs write their manifest to different locations depending on the
# training platform:
# HyperPod: <output>/<job>/manifest.json
# Serverless: <output>/<job>/output/output/manifest.json
# Serverful: <output>/<job>/output/output.tar.gz (manifest is inside)
# Try each in turn and surface every failure if none resolve, so the real
# cause is not masked by a misleading message from the last attempt.
hyperpod_manifest_uri = build_nova_hyperpod_manifest_s3_uri(
s3_output_path, training_job_name
)
serverless_manifest_uri = build_nova_manifest_s3_uri(s3_output_path, training_job_name)
tar_gz_uri = build_nova_output_tar_gz_s3_uri(s3_output_path, training_job_name)

attempts = [
("HyperPod manifest.json", hyperpod_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverless manifest.json", serverless_manifest_uri, read_nova_checkpoint_uri_from_manifest),
("serverful output.tar.gz", tar_gz_uri, _read_checkpoint_uri_from_tar_gz),
]

errors = []
for label, uri, reader in attempts:
try:
return reader(s3_client, uri)
except Exception as error: # noqa: PERF203 - each attempt may fail independently
errors.append(f"{label} at {uri} failed: {error}")

raise ValueError(
"Could not resolve the Nova checkpoint URI from any known output layout. "
+ " ".join(errors)
)
198 changes: 198 additions & 0 deletions sagemaker-core/tests/unit/test_training_utils.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,198 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Unit tests for Nova manifest/checkpoint helpers in training/utils.py."""
import io
import json
import tarfile

import pytest
from unittest.mock import Mock

from sagemaker.core.training.utils import (
build_nova_hyperpod_manifest_s3_uri,
build_nova_manifest_s3_uri,
build_nova_output_tar_gz_s3_uri,
read_nova_checkpoint_uri_from_manifest,
resolve_nova_checkpoint_uri,
)

CHECKPOINT_URI = "s3://bucket/ckpt/step_100"


def test_build_nova_manifest_s3_uri():
result = build_nova_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_manifest_s3_uri_strips_trailing_slash():
assert build_nova_manifest_s3_uri(
"s3://bucket/output//", "my-job"
) == "s3://bucket/output/my-job/output/output/manifest.json"


def test_build_nova_hyperpod_manifest_s3_uri():
result = build_nova_hyperpod_manifest_s3_uri("s3://bucket/output/", "my-job")
assert result == "s3://bucket/output/my-job/manifest.json"


def test_build_nova_output_tar_gz_s3_uri():
result = build_nova_output_tar_gz_s3_uri("s3://bucket/output", "my-job")
assert result == "s3://bucket/output/my-job/output/output.tar.gz"


def _s3_client_returning(body_bytes):
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
body = Mock()
body.read.return_value = body_bytes
client.get_object.return_value = {"Body": body}
return client


def test_read_manifest_returns_checkpoint_uri():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = read_nova_checkpoint_uri_from_manifest(
client, "s3://bucket/output/my-job/output/output/manifest.json"
)
assert result == CHECKPOINT_URI
client.get_object.assert_called_once_with(
Bucket="bucket", Key="output/my-job/output/output/manifest.json"
)


def test_read_manifest_missing_key_raises():
client = _s3_client_returning(json.dumps({"other": "value"}).encode("utf-8"))
with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_not_found_raises():
client = Mock()
client.exceptions = Mock()
client.exceptions.NoSuchKey = type("NoSuchKey", (Exception,), {})
client.get_object.side_effect = client.exceptions.NoSuchKey()
with pytest.raises(ValueError, match="manifest.json not found"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def test_read_manifest_invalid_json_raises():
client = _s3_client_returning(b"not-json")
with pytest.raises(ValueError, match="Failed to parse manifest.json"):
read_nova_checkpoint_uri_from_manifest(client, "s3://bucket/manifest.json")


def _make_tar_gz_with_manifest(manifest_dict):
content = json.dumps(manifest_dict).encode("utf-8")
buf = io.BytesIO()
with tarfile.open(fileobj=buf, mode="w:gz") as tar:
info = tarfile.TarInfo(name="manifest.json")
info.size = len(content)
tar.addfile(info, io.BytesIO(content))
return buf.getvalue()


def test_resolve_checkpoint_uri_from_raw_manifest():
client = _s3_client_returning(
json.dumps({"checkpoint_s3_bucket": CHECKPOINT_URI}).encode("utf-8")
)
result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_hyperpod_layout():
"""HyperPod writes the manifest at <output>/<job>/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
hyperpod_key = "output/my-job/manifest.json"

def get_object(Bucket, Key):
if Key != hyperpod_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_from_serverless_layout():
"""Serverless writes the manifest at <output>/<job>/output/output/manifest.json."""
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
serverless_key = "output/my-job/output/output/manifest.json"

def get_object(Bucket, Key):
if Key != serverless_key:
raise no_such_key()
body = Mock()
body.read.return_value = json.dumps(
{"checkpoint_s3_bucket": CHECKPOINT_URI}
).encode("utf-8")
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_falls_back_to_tar_gz():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"checkpoint_s3_bucket": CHECKPOINT_URI})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

result = resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
assert result == CHECKPOINT_URI


def test_resolve_checkpoint_uri_raises_when_both_sources_fail():
client = Mock()
client.exceptions = Mock()
no_such_key = type("NoSuchKey", (Exception,), {})
client.exceptions.NoSuchKey = no_such_key
tar_bytes = _make_tar_gz_with_manifest({"other": "value"})

def get_object(Bucket, Key):
if Key.endswith("manifest.json"):
raise no_such_key()
body = Mock()
body.read.return_value = tar_bytes
return {"Body": body}

client.get_object.side_effect = get_object

with pytest.raises(ValueError, match="checkpoint_s3_bucket"):
resolve_nova_checkpoint_uri(client, "s3://bucket/output/", "my-job")
Loading