diff --git a/sagemaker-core/src/sagemaker/core/image_uri_config/pytorch.json b/sagemaker-core/src/sagemaker/core/image_uri_config/pytorch.json index ab8046dcff..5dfa2122bf 100644 --- a/sagemaker-core/src/sagemaker/core/image_uri_config/pytorch.json +++ b/sagemaker-core/src/sagemaker/core/image_uri_config/pytorch.json @@ -1360,7 +1360,8 @@ "2.1": "2.1.0", "2.2": "2.2.1", "2.3": "2.3.0", - "2.4": "2.4.0" + "2.4": "2.4.0", + "2.6": "2.6.0" }, "versions": { "1.12.1": { @@ -1707,6 +1708,55 @@ "us-west-2": "763104351884" }, "repository": "pytorch-inference-graviton" + }, + "2.6.0": { + "container_version": { + "cpu": "ubuntu22.04" + }, + "py_versions": [ + "py312" + ], + "registries": { + "af-south-1": "626614931356", + "ap-east-1": "871362719292", + "ap-east-2": "975050140332", + "ap-northeast-1": "763104351884", + "ap-northeast-2": "763104351884", + "ap-northeast-3": "364406365360", + "ap-south-1": "763104351884", + "ap-south-2": "772153158452", + "ap-southeast-1": "763104351884", + "ap-southeast-2": "763104351884", + "ap-southeast-3": "907027046896", + "ap-southeast-4": "457447274322", + "ap-southeast-5": "550225433462", + "ap-southeast-6": "633930458069", + "ap-southeast-7": "590183813437", + "ca-central-1": "763104351884", + "ca-west-1": "204538143572", + "cn-north-1": "727897471807", + "cn-northwest-1": "727897471807", + "eu-central-1": "763104351884", + "eu-central-2": "380420809688", + "eu-north-1": "763104351884", + "eu-south-1": "692866216735", + "eu-south-2": "503227376785", + "eu-west-1": "763104351884", + "eu-west-2": "763104351884", + "eu-west-3": "763104351884", + "il-central-1": "780543022126", + "me-central-1": "914824155844", + "me-south-1": "217643126080", + "mx-central-1": "637423239942", + "sa-east-1": "763104351884", + "us-east-1": "763104351884", + "us-east-2": "763104351884", + "us-gov-east-1": "446045086412", + "us-gov-west-1": "442386744353", + "us-west-1": "763104351884", + "us-west-2": "763104351884" + }, + "repository": "pytorch-inference-arm64" } } }, diff --git a/sagemaker-core/src/sagemaker/core/image_uris.py b/sagemaker-core/src/sagemaker/core/image_uris.py index f3c063c34f..5d83aa297a 100644 --- a/sagemaker-core/src/sagemaker/core/image_uris.py +++ b/sagemaker-core/src/sagemaker/core/image_uris.py @@ -278,7 +278,7 @@ def retrieve( else: tag_prefix = version_config.get("tag_prefix", version) - if repo == f"{framework}-inference-graviton": + if repo in (f"{framework}-inference-graviton", f"{framework}-inference-arm64"): container_version = f"{container_version}-sagemaker" # Some images encode the accelerator directly in the tag (e.g. the amzn2023 diff --git a/sagemaker-core/tests/unit/image_uris/test_graviton.py b/sagemaker-core/tests/unit/image_uris/test_graviton.py index e498c8367a..800f057e0b 100644 --- a/sagemaker-core/tests/unit/image_uris/test_graviton.py +++ b/sagemaker-core/tests/unit/image_uris/test_graviton.py @@ -31,7 +31,13 @@ def _test_graviton_framework_uris( - framework, version, py_version, account, region, container_version="ubuntu20.04-sagemaker" + framework, + version, + py_version, + account, + region, + container_version="ubuntu20.04-sagemaker", + repository=None, ): for instance_type in GRAVITON_INSTANCE_TYPES: uri = image_uris.retrieve(framework, region, instance_type=instance_type, version=version) @@ -42,6 +48,7 @@ def _test_graviton_framework_uris( account, region=region, container_version=container_version, + repository=repository, ) assert expected == uri @@ -57,6 +64,9 @@ def test_graviton_framework_uris(load_config_and_file_name, scope): for version in VERSIONS: ACCOUNTS = config[scope]["versions"][version]["registries"] py_versions = config[scope]["versions"][version]["py_versions"] + repository = config[scope]["versions"][version].get( + "repository", "{}-inference-graviton".format(framework) + ) container_version = ( config[scope]["versions"][version].get("container_version", {}).get("cpu", None) ) @@ -66,11 +76,22 @@ def test_graviton_framework_uris(load_config_and_file_name, scope): for region in ACCOUNTS.keys(): if container_version: _test_graviton_framework_uris( - framework, version, py_version, ACCOUNTS[region], region, container_version + framework, + version, + py_version, + ACCOUNTS[region], + region, + container_version, + repository, ) else: _test_graviton_framework_uris( - framework, version, py_version, ACCOUNTS[region], region + framework, + version, + py_version, + ACCOUNTS[region], + region, + repository=repository, ) @@ -201,10 +222,10 @@ def test_graviton_sklearn_image_scope_specified_x86_instance(graviton_sklearn_un def _expected_graviton_framework_uri( - framework, version, py_version, account, region, container_version + framework, version, py_version, account, region, container_version, repository=None ): return expected_uris.graviton_framework_uri( - "{}-inference-graviton".format(framework), + repository or "{}-inference-graviton".format(framework), fw_version=version, py_version=py_version, account=account,