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
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
"""

import json
import logging
import boto3
from typing import Union, Optional, Dict, Any
from dataclasses import dataclass
Expand All @@ -16,6 +17,8 @@
from sagemaker.train.constants import get_sagemaker_hub_name
from sagemaker.core.utils.utils import Unassigned

_logger = logging.getLogger(__name__)


class _ModelType(Enum):
"""Internal enum for model type classification."""
Expand Down Expand Up @@ -368,7 +371,10 @@ def _resolve_model_package_object(self, model_package: 'ModelPackage') -> _Model
# integ-test hub) resolve correctly. Public-hub content is
# account-less ("aws"); private-hub content lives under the
# model package's own account.
hub_name = get_sagemaker_hub_name()
#
hub_name = self._resolve_base_model_hub(
hub_content_name, hub_content_version, region
)
hub_account = "aws" if hub_name == "SageMakerPublicHub" else account
base_model_arn = f"arn:aws:sagemaker:{region}:{hub_account}:hub-content/{hub_name}/Model/{hub_content_name}/{hub_content_version}"

Expand Down Expand Up @@ -464,6 +470,52 @@ def _validate_model_package_arn(self, arn: str) -> bool:
)
return True

def _resolve_base_model_hub(
self, hub_content_name: str, hub_content_version: str, region: str
) -> str:
"""Pick the hub that backs the base model's hub content.
Returns the hub from ``SAGEMAKER_HUB_NAME`` (defaults to
``SageMakerPublicHub``). When a non-public hub is configured but does
not contain the base model, returns ``SageMakerPublicHub`` instead,
since base models are always published there.
Args:
hub_content_name: Base model hub content name.
hub_content_version: Base model hub content version.
region: AWS region of the model package.
Returns:
The hub name whose content should back the base model ARN.
"""
hub_name = get_sagemaker_hub_name()
if hub_name == "SageMakerPublicHub":
return hub_name

from sagemaker.core.resources import HubContent

try:
session = self._get_session()
HubContent.get(
hub_name=hub_name,
hub_content_type="Model",
hub_content_name=hub_content_name,
hub_content_version=hub_content_version,
session=session.boto_session,
region=region,
)
return hub_name
except Exception as e:
_logger.info(
"Base model '%s' (v%s) not found in hub '%s' (%s); "
"falling back to SageMakerPublicHub.",
hub_content_name,
hub_content_version,
hub_name,
e,
)
return "SageMakerPublicHub"

def _get_session(self):
"""
Get or create SageMaker session.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -440,19 +440,23 @@ def test_resolve_arn_construct_hub_content_arn(self, mock_validate, mock_get_ses
assert result.base_model_name == "base-model"
assert result.hub_content_name == "base-model"

@patch('sagemaker.core.resources.HubContent')
@patch('sagemaker.core.resources.ModelPackage')
@patch('sagemaker.train.common_utils.model_resolution._ModelResolver._get_session')
@patch('sagemaker.train.common_utils.model_resolution._ModelResolver._validate_model_package_arn')
def test_resolve_arn_construct_hub_content_arn_private_hub(self, mock_validate, mock_get_session, mock_model_package_class):
"""When SAGEMAKER_HUB_NAME points at a private hub, the reconstructed
base-model ARN targets that hub under the model package's own account
(not the account-less public hub)."""
def test_resolve_arn_construct_hub_content_arn_private_hub(self, mock_validate, mock_get_session, mock_model_package_class, mock_hub_content_class):
"""When SAGEMAKER_HUB_NAME points at a private hub that DOES contain the
base model, the reconstructed base-model ARN targets that hub under the
model package's own account (not the account-less public hub)."""
arn = "arn:aws:sagemaker:us-west-2:123456789012:model-package/my-model/1"

mock_session = MagicMock()
mock_session.boto_session.region_name = 'us-west-2'
mock_get_session.return_value = mock_session

# Private hub contains the base model, so verification succeeds.
mock_hub_content_class.get.return_value = MagicMock()

# Mock ModelPackage without hub_content_arn (needs to be constructed)
mock_package = MagicMock()
mock_package.model_package_arn = arn
Expand All @@ -478,6 +482,54 @@ def test_resolve_arn_construct_hub_content_arn_private_hub(self, mock_validate,
assert result.base_model_arn == expected_arn
assert result.base_model_name == "mock-oss-test"
assert result.hub_content_name == "mock-oss-test"

@patch('sagemaker.core.resources.HubContent')
@patch('sagemaker.core.resources.ModelPackage')
@patch('sagemaker.train.common_utils.model_resolution._ModelResolver._get_session')
@patch('sagemaker.train.common_utils.model_resolution._ModelResolver._validate_model_package_arn')
def test_resolve_arn_construct_hub_content_arn_private_hub_fallback_public(self, mock_validate, mock_get_session, mock_model_package_class, mock_hub_content_class):
"""When SAGEMAKER_HUB_NAME points at a private hub that does NOT contain
the base model (e.g. it never mirrored it or was cleaned up), the
reconstructed base-model ARN falls back to the account-less public hub.

Regression test: without this fallback the server-side CreateTrainingJob
fails with 'Hub content ... does not exist' during evaluation."""
arn = "arn:aws:sagemaker:us-west-2:123456789012:model-package/my-model/1"

mock_session = MagicMock()
mock_session.boto_session.region_name = 'us-west-2'
mock_get_session.return_value = mock_session

# Private hub does NOT contain the base model -> verification raises.
mock_hub_content_class.get.side_effect = Exception(
"Hub content with name mock-oss-test does not exist."
)

# Mock ModelPackage without hub_content_arn (needs to be constructed)
mock_package = MagicMock()
mock_package.model_package_arn = arn

mock_container = MagicMock()
mock_base_model = MagicMock()
mock_base_model.hub_content_name = 'mock-oss-test'
mock_base_model.hub_content_version = '0.0.1'
mock_base_model.hub_content_arn = None # Not provided, needs construction
mock_container.base_model = mock_base_model

mock_package.inference_specification = MagicMock()
mock_package.inference_specification.containers = [mock_container]

mock_model_package_class.get.return_value = mock_package

resolver = _ModelResolver()
with patch.dict(os.environ, {"SAGEMAKER_HUB_NAME": "sdktest"}):
result = resolver._resolve_model_package_arn(arn)

# Fell back to the account-less public hub since the private hub lacks it.
expected_arn = "arn:aws:sagemaker:us-west-2:aws:hub-content/SageMakerPublicHub/Model/mock-oss-test/0.0.1"
assert result.base_model_arn == expected_arn
assert result.base_model_name == "mock-oss-test"
assert result.hub_content_name == "mock-oss-test"

@patch('sagemaker.core.resources.ModelPackage')
@patch('sagemaker.train.common_utils.model_resolution._ModelResolver._get_session')
Expand Down
Loading