From c62cbb8d82f303164ab70de0ac3700d85f554993 Mon Sep 17 00:00:00 2001 From: Amarjeet LNU Date: Thu, 16 Jul 2026 18:12:39 -0700 Subject: [PATCH] fix: resolve MTRL eval base-model ARN against the configured hub When evaluating a fine-tuned model via attach()/ModelPackage, the base model's hub-content ARN is reconstructed from the model package's BaseModel metadata (HubContentName + HubContentVersion) because the backend does not populate HubContentArn. That reconstruction was hardcoded to SageMakerPublicHub with an account-less ("aws") owner. Models customized against a private/custom hub (SAGEMAKER_HUB_NAME) then resolve to a public-hub ARN that does not exist, and the evaluation pipeline's CreateJob fails server-side with: ResourceNotFound: Hub content with name does not exist Honor get_sagemaker_hub_name() when reconstructing the ARN, and use the model package's own account for private hubs (public-hub content stays account-less). The string/JumpStart-ID path already honored the hub via _resolve_jumpstart_model; this aligns the model-package path with it. Add a private-hub unit test and pin the existing test to the default hub. --- .../train/common_utils/model_resolution.py | 16 +++-- .../common_utils/test_model_resolution.py | 63 ++++++++++++++++--- 2 files changed, 66 insertions(+), 13 deletions(-) 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 c1d8cdc91e..418e6b5110 100644 --- a/sagemaker-train/src/sagemaker/train/common_utils/model_resolution.py +++ b/sagemaker-train/src/sagemaker/train/common_utils/model_resolution.py @@ -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: 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 cf1805cd50..8a2ee3be22 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 @@ -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')