Skip to content
Open
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 @@ -357,12 +357,20 @@ def _resolve_model_package_object(self, model_package: 'ModelPackage') -> _Model
model_pkg_arn = getattr(model_package, 'model_package_arn', None)

if hub_content_name and hub_content_version and model_pkg_arn:
# Extract region from model package ARN
# Extract region and account from model package ARN
arn_parts = model_pkg_arn.split(':')
if len(arn_parts) >= 4:
if len(arn_parts) >= 5:
region = arn_parts[3]
# Base model always lives in SageMakerPublicHub (SAGEMAKER_HUB_NAME is for training recipes only)
base_model_arn = f"arn:aws:sagemaker:{region}:aws:hub-content/SageMakerPublicHub/Model/{hub_content_name}/{hub_content_version}"
account = arn_parts[4]
# Reconstruct the base-model hub-content ARN in the hub the
# model was customized against. Defaults to SageMakerPublicHub
# but honors SAGEMAKER_HUB_NAME so private/custom hubs (e.g. an
# 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_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}"

# If we couldn't extract or construct base model ARN, this is not a supported model package
if not base_model_arn:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -401,38 +401,83 @@ def test_resolve_arn_success(self, mock_validate, mock_get_session, mock_model_p
@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(self, mock_validate, mock_get_session, mock_model_package_class):
"""Test ARN resolution when HubContentArn needs to be constructed."""
"""Test ARN resolution when HubContentArn needs to be constructed.

With no SAGEMAKER_HUB_NAME override, the base model is assumed to live
in the account-less SageMakerPublicHub.
"""
arn = "arn:aws:sagemaker:us-west-2:123456789012:model-package/my-model/1"

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

# 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 = 'base-model'
mock_base_model.hub_content_version = '1.0'
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()
result = resolver._resolve_model_package_arn(arn)

with patch.dict(os.environ, {}, clear=False):
os.environ.pop("SAGEMAKER_HUB_NAME", None)
result = resolver._resolve_model_package_arn(arn)

# Should construct ARN from region and hub content name/version
expected_arn = "arn:aws:sagemaker:us-west-2:aws:hub-content/SageMakerPublicHub/Model/base-model/1.0"
assert result.base_model_arn == expected_arn
assert result.base_model_name == "base-model"
assert result.hub_content_name == "base-model"

@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)."""
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

# 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)

# Private hub: uses the model package's account (123456789012), not "aws"
expected_arn = "arn:aws:sagemaker:us-west-2:123456789012:hub-content/sdktest/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