Skip to content

Commit 87dc675

Browse files
authored
fix: resolve MTRL eval base-model ARN against the configured hub (#6040)
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 <model> 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.
1 parent e9aa599 commit 87dc675

2 files changed

Lines changed: 66 additions & 13 deletions

File tree

sagemaker-train/src/sagemaker/train/common_utils/model_resolution.py

Lines changed: 12 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -357,12 +357,20 @@ def _resolve_model_package_object(self, model_package: 'ModelPackage') -> _Model
357357
model_pkg_arn = getattr(model_package, 'model_package_arn', None)
358358

359359
if hub_content_name and hub_content_version and model_pkg_arn:
360-
# Extract region from model package ARN
360+
# Extract region and account from model package ARN
361361
arn_parts = model_pkg_arn.split(':')
362-
if len(arn_parts) >= 4:
362+
if len(arn_parts) >= 5:
363363
region = arn_parts[3]
364-
# Base model always lives in SageMakerPublicHub (SAGEMAKER_HUB_NAME is for training recipes only)
365-
base_model_arn = f"arn:aws:sagemaker:{region}:aws:hub-content/SageMakerPublicHub/Model/{hub_content_name}/{hub_content_version}"
364+
account = arn_parts[4]
365+
# Reconstruct the base-model hub-content ARN in the hub the
366+
# model was customized against. Defaults to SageMakerPublicHub
367+
# but honors SAGEMAKER_HUB_NAME so private/custom hubs (e.g. an
368+
# integ-test hub) resolve correctly. Public-hub content is
369+
# account-less ("aws"); private-hub content lives under the
370+
# model package's own account.
371+
hub_name = get_sagemaker_hub_name()
372+
hub_account = "aws" if hub_name == "SageMakerPublicHub" else account
373+
base_model_arn = f"arn:aws:sagemaker:{region}:{hub_account}:hub-content/{hub_name}/Model/{hub_content_name}/{hub_content_version}"
366374

367375
# If we couldn't extract or construct base model ARN, this is not a supported model package
368376
if not base_model_arn:

sagemaker-train/tests/unit/train/common_utils/test_model_resolution.py

Lines changed: 54 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -401,38 +401,83 @@ def test_resolve_arn_success(self, mock_validate, mock_get_session, mock_model_p
401401
@patch('sagemaker.train.common_utils.model_resolution._ModelResolver._get_session')
402402
@patch('sagemaker.train.common_utils.model_resolution._ModelResolver._validate_model_package_arn')
403403
def test_resolve_arn_construct_hub_content_arn(self, mock_validate, mock_get_session, mock_model_package_class):
404-
"""Test ARN resolution when HubContentArn needs to be constructed."""
404+
"""Test ARN resolution when HubContentArn needs to be constructed.
405+
406+
With no SAGEMAKER_HUB_NAME override, the base model is assumed to live
407+
in the account-less SageMakerPublicHub.
408+
"""
405409
arn = "arn:aws:sagemaker:us-west-2:123456789012:model-package/my-model/1"
406-
410+
407411
# Mock session
408412
mock_session = MagicMock()
409413
mock_session.boto_session.region_name = 'us-west-2'
410414
mock_get_session.return_value = mock_session
411-
415+
412416
# Mock ModelPackage without hub_content_arn (needs to be constructed)
413417
mock_package = MagicMock()
414418
mock_package.model_package_arn = arn
415-
419+
416420
mock_container = MagicMock()
417421
mock_base_model = MagicMock()
418422
mock_base_model.hub_content_name = 'base-model'
419423
mock_base_model.hub_content_version = '1.0'
420424
mock_base_model.hub_content_arn = None # Not provided, needs construction
421425
mock_container.base_model = mock_base_model
422-
426+
423427
mock_package.inference_specification = MagicMock()
424428
mock_package.inference_specification.containers = [mock_container]
425-
429+
426430
mock_model_package_class.get.return_value = mock_package
427-
431+
428432
resolver = _ModelResolver()
429-
result = resolver._resolve_model_package_arn(arn)
430-
433+
with patch.dict(os.environ, {}, clear=False):
434+
os.environ.pop("SAGEMAKER_HUB_NAME", None)
435+
result = resolver._resolve_model_package_arn(arn)
436+
431437
# Should construct ARN from region and hub content name/version
432438
expected_arn = "arn:aws:sagemaker:us-west-2:aws:hub-content/SageMakerPublicHub/Model/base-model/1.0"
433439
assert result.base_model_arn == expected_arn
434440
assert result.base_model_name == "base-model"
435441
assert result.hub_content_name == "base-model"
442+
443+
@patch('sagemaker.core.resources.ModelPackage')
444+
@patch('sagemaker.train.common_utils.model_resolution._ModelResolver._get_session')
445+
@patch('sagemaker.train.common_utils.model_resolution._ModelResolver._validate_model_package_arn')
446+
def test_resolve_arn_construct_hub_content_arn_private_hub(self, mock_validate, mock_get_session, mock_model_package_class):
447+
"""When SAGEMAKER_HUB_NAME points at a private hub, the reconstructed
448+
base-model ARN targets that hub under the model package's own account
449+
(not the account-less public hub)."""
450+
arn = "arn:aws:sagemaker:us-west-2:123456789012:model-package/my-model/1"
451+
452+
mock_session = MagicMock()
453+
mock_session.boto_session.region_name = 'us-west-2'
454+
mock_get_session.return_value = mock_session
455+
456+
# Mock ModelPackage without hub_content_arn (needs to be constructed)
457+
mock_package = MagicMock()
458+
mock_package.model_package_arn = arn
459+
460+
mock_container = MagicMock()
461+
mock_base_model = MagicMock()
462+
mock_base_model.hub_content_name = 'mock-oss-test'
463+
mock_base_model.hub_content_version = '0.0.1'
464+
mock_base_model.hub_content_arn = None # Not provided, needs construction
465+
mock_container.base_model = mock_base_model
466+
467+
mock_package.inference_specification = MagicMock()
468+
mock_package.inference_specification.containers = [mock_container]
469+
470+
mock_model_package_class.get.return_value = mock_package
471+
472+
resolver = _ModelResolver()
473+
with patch.dict(os.environ, {"SAGEMAKER_HUB_NAME": "sdktest"}):
474+
result = resolver._resolve_model_package_arn(arn)
475+
476+
# Private hub: uses the model package's account (123456789012), not "aws"
477+
expected_arn = "arn:aws:sagemaker:us-west-2:123456789012:hub-content/sdktest/Model/mock-oss-test/0.0.1"
478+
assert result.base_model_arn == expected_arn
479+
assert result.base_model_name == "mock-oss-test"
480+
assert result.hub_content_name == "mock-oss-test"
436481

437482
@patch('sagemaker.core.resources.ModelPackage')
438483
@patch('sagemaker.train.common_utils.model_resolution._ModelResolver._get_session')

0 commit comments

Comments
 (0)