diff --git a/sagemaker-core/src/sagemaker/core/common_utils.py b/sagemaker-core/src/sagemaker/core/common_utils.py index 0c8025174c..6077b04a4d 100644 --- a/sagemaker-core/src/sagemaker/core/common_utils.py +++ b/sagemaker-core/src/sagemaker/core/common_utils.py @@ -2205,19 +2205,24 @@ def camel_to_snake(camel_case_string: str) -> str: def walk_and_apply_json( - json_obj: Dict[Any, Any], apply, stop_keys: Optional[List[str]] = ["metrics"] + json_obj: Dict[Any, Any], + apply, + stop_keys: Optional[List[str]] = ["metrics", "environment_variables"], ) -> Dict[Any, Any]: """Recursively walks a json object and applies a given function to the keys. stop_keys (Optional[list[str]]): List of field keys that should stop the application function. Any children of these keys will not have the application function applied to them. + A key stops the walk if either its original or its converted form is in stop_keys, so + the same list works for camel_to_snake and snake_to_upper_camel passes. Environment + variable names are user facing values stored as keys and must never be converted. """ def _walk_and_apply_json(json_obj, new): if isinstance(json_obj, dict) and isinstance(new, dict): for key, value in json_obj.items(): new_key = apply(key) - if (stop_keys and new_key not in stop_keys) or stop_keys is None: + if stop_keys is None or (key not in stop_keys and new_key not in stop_keys): if isinstance(value, dict): new[new_key] = {} _walk_and_apply_json(value, new=new[new_key]) diff --git a/sagemaker-core/src/sagemaker/core/jumpstart/hub/parser_utils.py b/sagemaker-core/src/sagemaker/core/jumpstart/hub/parser_utils.py index 0983122d09..292388af00 100644 --- a/sagemaker-core/src/sagemaker/core/jumpstart/hub/parser_utils.py +++ b/sagemaker-core/src/sagemaker/core/jumpstart/hub/parser_utils.py @@ -33,19 +33,24 @@ def snake_to_upper_camel(snake_case_string: str) -> str: def walk_and_apply_json( - json_obj: Dict[Any, Any], apply, stop_keys: Optional[List[str]] = ["metrics"] + json_obj: Dict[Any, Any], + apply, + stop_keys: Optional[List[str]] = ["metrics", "environment_variables"], ) -> Dict[Any, Any]: """Recursively walks a json object and applies a given function to the keys. stop_keys (Optional[list[str]]): List of field keys that should stop the application function. Any children of these keys will not have the application function applied to them. + A key stops the walk if either its original or its converted form is in stop_keys, so + the same list works for camel_to_snake and snake_to_upper_camel passes. Environment + variable names are user facing values stored as keys and must never be converted. """ def _walk_and_apply_json(json_obj, new): if isinstance(json_obj, dict) and isinstance(new, dict): for key, value in json_obj.items(): new_key = apply(key) - if (stop_keys and new_key not in stop_keys) or stop_keys is None: + if stop_keys is None or (key not in stop_keys and new_key not in stop_keys): if isinstance(value, dict): new[new_key] = {} _walk_and_apply_json(value, new=new[new_key]) diff --git a/sagemaker-core/tests/unit/jumpstart/artifacts/test_environment_variables.py b/sagemaker-core/tests/unit/jumpstart/artifacts/test_environment_variables.py new file mode 100644 index 0000000000..72895abc69 --- /dev/null +++ b/sagemaker-core/tests/unit/jumpstart/artifacts/test_environment_variables.py @@ -0,0 +1,82 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"). You +# may not use this file except in compliance with the License. A copy of +# the License is located at +# +# http://aws.amazon.com/apache2.0/ +# +# or in the "license" file accompanying this file. This file is +# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF +# ANY KIND, either express or implied. See the License for the specific +# language governing permissions and limitations under the License. +"""Unit tests for JumpStart environment variable retrieval.""" +from __future__ import absolute_import + +import datetime +import json +import os +from unittest.mock import patch + +import pytest + +from sagemaker.core.jumpstart.artifacts.environment_variables import ( + _retrieve_default_environment_variables, +) +from sagemaker.core.jumpstart.hub.interfaces import DescribeHubContentResponse +from sagemaker.core.jumpstart.hub.parsers import ( + make_model_specs_from_describe_hub_content_response, +) + +HUB_ARN = "arn:aws:sagemaker:us-east-1:123456789012:hub/my-private-hub" + + +@pytest.fixture +def private_hub_model_specs(): + """Model specs parsed from a private hub document, the way the JumpStart cache builds them.""" + path = os.path.join(os.path.dirname(__file__), "..", "hub_content_document.json") + with open(path, "r") as f: + hub_content_document = json.load(f) + response = DescribeHubContentResponse( + { + "CreationTime": datetime.datetime(2024, 1, 1), + "DocumentSchemaVersion": "2.0.0", + "HubArn": HUB_ARN, + "HubContentArn": ( + "arn:aws:sagemaker:us-east-1:123456789012:hub-content/" + "my-private-hub/Model/meta-textgeneration-llama-2-13b-f/1.0.0" + ), + "HubContentName": "meta-textgeneration-llama-2-13b-f", + "HubContentType": "Model", + "HubContentVersion": "1.0.0", + "HubContentStatus": "Available", + "HubName": "my-private-hub", + "HubContentDocument": json.dumps(hub_content_document), + } + ) + return make_model_specs_from_describe_hub_content_response(response) + + +@patch( + "sagemaker.core.jumpstart.artifacts.environment_variables.verify_model_region_and_return_specs" +) +def test_private_hub_instance_specific_environment_variable_overrides_default( + mock_verify_model_region_and_return_specs, private_hub_model_specs +): + """An instance specific value must replace the model default rather than being added + under a mangled second key, which is what ended up in the container environment before.""" + mock_verify_model_region_and_return_specs.return_value = private_hub_model_specs + + environment_variables = _retrieve_default_environment_variables( + model_id="meta-textgeneration-llama-2-13b-f", + model_version="1.0.0", + hub_arn=HUB_ARN, + region="us-east-1", + instance_type="ml.g5.48xlarge", + sagemaker_session=None, + ) + + gpu_keys = [key for key in environment_variables if key.replace("_", "").lower() == "smnumgpus"] + assert gpu_keys == ["SM_NUM_GPUS"] + assert environment_variables["SM_NUM_GPUS"] == "8" + assert mock_verify_model_region_and_return_specs.call_args.kwargs["hub_arn"] == HUB_ARN diff --git a/sagemaker-core/tests/unit/jumpstart/hub/test_parser_utils.py b/sagemaker-core/tests/unit/jumpstart/hub/test_parser_utils.py new file mode 100644 index 0000000000..08bbc60ea0 --- /dev/null +++ b/sagemaker-core/tests/unit/jumpstart/hub/test_parser_utils.py @@ -0,0 +1,89 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"). You +# may not use this file except in compliance with the License. A copy of +# the License is located at +# +# http://aws.amazon.com/apache2.0/ +# +# or in the "license" file accompanying this file. This file is +# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF +# ANY KIND, either express or implied. See the License for the specific +# language governing permissions and limitations under the License. +from __future__ import absolute_import + +from sagemaker.core.jumpstart.hub.parser_utils import ( + camel_to_snake, + snake_to_upper_camel, + walk_and_apply_json, +) + + +VARIANTS = { + "Variants": { + "g5": { + "Properties": { + "ImageUri": "image", + "EnvironmentVariables": { + "SM_VLLM_MAX_MODEL_LEN": "131072", + "HF_HUB_OFFLINE": "1", + }, + "Metrics": [{"Name": "loss", "Regex": "loss=(.*)"}], + } + } + } +} + + +def test_walk_and_apply_json_camel_to_snake_preserves_environment_variable_names(): + result = walk_and_apply_json(VARIANTS, camel_to_snake) + + properties = result["variants"]["g5"]["properties"] + assert properties["image_uri"] == "image" + assert properties["environment_variables"] == { + "SM_VLLM_MAX_MODEL_LEN": "131072", + "HF_HUB_OFFLINE": "1", + } + assert properties["metrics"] == [{"Name": "loss", "Regex": "loss=(.*)"}] + + +def test_walk_and_apply_json_snake_to_upper_camel_preserves_environment_variable_names(): + snake = walk_and_apply_json(VARIANTS, camel_to_snake) + + result = walk_and_apply_json(snake, snake_to_upper_camel) + + properties = result["Variants"]["G5"]["Properties"] + assert properties["ImageUri"] == "image" + assert properties["EnvironmentVariables"] == { + "SM_VLLM_MAX_MODEL_LEN": "131072", + "HF_HUB_OFFLINE": "1", + } + + +def test_walk_and_apply_json_round_trip_preserves_environment_variable_names(): + """The private hub pipeline converts the same document several times in both directions.""" + result = walk_and_apply_json(VARIANTS, camel_to_snake) + result = walk_and_apply_json(result, snake_to_upper_camel) + result = walk_and_apply_json(result, camel_to_snake) + result = walk_and_apply_json(result, camel_to_snake) + + assert result["variants"]["g5"]["properties"]["environment_variables"] == { + "SM_VLLM_MAX_MODEL_LEN": "131072", + "HF_HUB_OFFLINE": "1", + } + + +def test_walk_and_apply_json_explicit_stop_keys_still_honored(): + result = walk_and_apply_json( + {"Outer": {"Inner": {"KeepMe": 1}}}, camel_to_snake, stop_keys=["inner"] + ) + + assert result == {"outer": {"inner": {"KeepMe": 1}}} + + +def test_walk_and_apply_json_no_stop_keys_converts_everything(): + result = walk_and_apply_json(VARIANTS, camel_to_snake, stop_keys=None) + + properties = result["variants"]["g5"]["properties"] + assert "SM_VLLM_MAX_MODEL_LEN" not in properties["environment_variables"] + assert properties["metrics"] == [{"name": "loss", "regex": "loss=(.*)"}] diff --git a/sagemaker-core/tests/unit/jumpstart/hub/test_parsers.py b/sagemaker-core/tests/unit/jumpstart/hub/test_parsers.py index 1eb30889a3..310aa418f2 100644 --- a/sagemaker-core/tests/unit/jumpstart/hub/test_parsers.py +++ b/sagemaker-core/tests/unit/jumpstart/hub/test_parsers.py @@ -11,6 +11,10 @@ # ANY KIND, either express or implied. See the License for the specific # language governing permissions and limitations under the License. +import copy +import datetime +import json +import os import pytest from unittest.mock import Mock, patch from sagemaker.core.jumpstart.hub.parsers import ( @@ -367,3 +371,87 @@ def test_make_model_specs_from_describe_hub_content_response_with_payloads(self) result = make_model_specs_from_describe_hub_content_response(response) assert result is not None + + def test_make_model_specs_preserves_instance_variant_environment_variable_names( + self, hub_content_document + ): + """Environment variable names under HostingInstanceTypeVariants must survive verbatim.""" + variants = hub_content_document["InferenceConfigComponents"]["tgi"][ + "HostingInstanceTypeVariants" + ]["Variants"] + variants["ml.g5.12xlarge"]["Properties"]["EnvironmentVariables"][ + "SM_VLLM_MAX_MODEL_LEN" + ] = "4096" + variants["g5"]["Properties"]["EnvironmentVariables"] = {"HF_HUB_OFFLINE": "1"} + + specs = make_model_specs_from_describe_hub_content_response( + _describe_hub_content_response(hub_content_document) + ) + + instance_type_variants = specs.hosting_instance_type_variants + assert instance_type_variants.variants["ml.g5.12xlarge"]["properties"][ + "environment_variables" + ] == {"SM_NUM_GPUS": "4", "SM_VLLM_MAX_MODEL_LEN": "4096"} + assert instance_type_variants.get_instance_specific_environment_variables( + "ml.g5.12xlarge" + ) == { + "HF_HUB_OFFLINE": "1", + "SM_NUM_GPUS": "4", + "SM_VLLM_MAX_MODEL_LEN": "4096", + } + + def test_make_model_specs_preserves_top_level_variant_environment_variable_names( + self, hub_content_document + ): + """Top level HostingInstanceTypeVariants go through one more conversion pass than + inference config components and must also keep environment variable names verbatim.""" + document = copy.deepcopy(hub_content_document) + for key in ("InferenceConfigs", "InferenceConfigComponents", "InferenceConfigRankings"): + document.pop(key, None) + document["HostingInstanceTypeVariants"] = { + "Variants": { + "g5": { + "Properties": { + "ImageUri": "image", + "EnvironmentVariables": {"SM_VLLM_MAX_MODEL_LEN": "131072"}, + } + } + } + } + + specs = make_model_specs_from_describe_hub_content_response( + _describe_hub_content_response(document) + ) + + instance_type_variants = specs.hosting_instance_type_variants + assert instance_type_variants.variants["g5"]["properties"]["image_uri"] == "image" + assert instance_type_variants.get_instance_specific_environment_variables( + "ml.g5.12xlarge" + ) == {"SM_VLLM_MAX_MODEL_LEN": "131072"} + + +def _describe_hub_content_response(hub_content_document): + return DescribeHubContentResponse( + { + "CreationTime": datetime.datetime(2024, 1, 1), + "DocumentSchemaVersion": "2.0.0", + "HubArn": "arn:aws:sagemaker:us-east-1:123456789012:hub/my-private-hub", + "HubContentArn": ( + "arn:aws:sagemaker:us-east-1:123456789012:hub-content/" + "my-private-hub/Model/meta-textgeneration-llama-2-13b-f/1.0.0" + ), + "HubContentName": "meta-textgeneration-llama-2-13b-f", + "HubContentType": "Model", + "HubContentVersion": "1.0.0", + "HubContentStatus": "Available", + "HubName": "my-private-hub", + "HubContentDocument": json.dumps(hub_content_document), + } + ) + + +@pytest.fixture +def hub_content_document(): + path = os.path.join(os.path.dirname(__file__), "..", "hub_content_document.json") + with open(path, "r") as f: + return json.load(f)