diff --git a/sagemaker-train/src/sagemaker/train/common_utils/model_resolution.py b/sagemaker-train/src/sagemaker/train/common_utils/model_resolution.py index 418e6b5110..f047767e1f 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/model_resolution.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/model_resolution.py @@ -7,6 +7,7 @@ """ import json +import logging import boto3 from typing import Union, Optional, Dict, Any from dataclasses import dataclass @@ -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.""" @@ -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}" @@ -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. diff --git a/sagemaker-train/tests/unit/train/common_utils/test_model_resolution.py b/sagemaker-train/tests/unit/train/common_utils/test_model_resolution.py index 8a2ee3be22..82d4e5d752 100644 --- a/sagemaker-train/tests/unit/train/common_utils/test_model_resolution.py +++ b/sagemaker-train/tests/unit/train/common_utils/test_model_resolution.py @@ -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 @@ -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')