34 changes: 32 additions & 2 deletions src/sagemaker/jumpstart/artifacts/environment_variables.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,10 +12,11 @@
# language governing permissions and limitations under the License.
"""This module contains functions for obtaining JumpStart environment variables."""
from __future__ import absolute_import
from typing import Dict, Optional
from typing import Callable, Dict, Optional, Set
from sagemaker.jumpstart.constants import (
DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
JUMPSTART_DEFAULT_REGION_NAME,
JUMPSTART_LOGGER,
SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY,
)
from sagemaker.jumpstart.enums import (
Expand DownExpand Up@@ -110,7 +111,9 @@ def _retrieve_default_environment_variables(

default_environment_variables.update(instance_specific_environment_variables)

gated_model_env_var: Optional[str] = _retrieve_gated_model_uri_env_var_value(
retrieve_gated_env_var_for_instance_type: Callable[
[str], Optional[str]
] = lambda instance_type: _retrieve_gated_model_uri_env_var_value(
model_id=model_id,
model_version=model_version,
region=region,
Expand All@@ -120,6 +123,33 @@ def _retrieve_default_environment_variables(
instance_type=instance_type,
)

gated_model_env_var: Optional[str] = retrieve_gated_env_var_for_instance_type(
instance_type
)

if gated_model_env_var is None and model_specs.is_gated_model():

possible_env_vars: Set[str] = {
retrieve_gated_env_var_for_instance_type(instance_type)
for instance_type in model_specs.supported_training_instance_types
}

# If all officially supported instance types have the same underlying artifact,
# we can use this artifact with high confidence that it'll succeed with
# an arbitrary instance.
if len(possible_env_vars) == 1:
gated_model_env_var = list(possible_env_vars)[0]

# If this model does not have 1 artifact for all supported instance types,
# we cannot determine which artifact to use for an arbitrary instance.
else:
log_msg = (
f"'{model_id}' does not support {instance_type} instance type"
" for training. Please use one of the following instance types: "
f"{', '.join(model_specs.supported_training_instance_types)}."
)
JUMPSTART_LOGGER.warning(log_msg)

if gated_model_env_var is not None:
default_environment_variables.update(
{SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY: gated_model_env_var}
Expand Down
21 changes: 21 additions & 0 deletions src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -62,6 +62,7 @@
)
from sagemaker.jumpstart.utils import (
add_jumpstart_model_id_version_tags,
get_eula_message,
update_dict_if_key_not_present,
resolve_estimator_sagemaker_config_field,
verify_model_region_and_return_specs,
Expand DownExpand Up@@ -597,6 +598,26 @@ def _add_env_to_kwargs(
value,
)

environment = getattr(kwargs, "environment", {}) or {}
if (
environment.get(SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY)
and str(environment.get("accept_eula", "")).lower() != "true"
):
model_specs = verify_model_region_and_return_specs(
model_id=kwargs.model_id,
version=kwargs.model_version,
region=kwargs.region,
scope=JumpStartScriptScope.TRAINING,
tolerate_deprecated_model=kwargs.tolerate_deprecated_model,
tolerate_vulnerable_model=kwargs.tolerate_vulnerable_model,
sagemaker_session=kwargs.sagemaker_session,
)
if model_specs.is_gated_model():

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Question: shouldn't this be the outer if statement here? And we would throw an error in all of the cases (the environment does not contain the special key, accept_eula is missing, accept_eula is false, etc.)?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'd rather make the conditions as specific as possible. Maybe we'll change how we launch gated models in the future.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

+1, this is odd, why bother retrieving the specs if you're not going to use them?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

fwiw, this won't involve another s3 call, it's just reading from memory

raise ValueError(
"Need to define ‘accept_eula'='true' within Environment. "
f"{get_eula_message(model_specs, kwargs.region)}"
)

return kwargs


Expand Down
4 changes: 4 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -963,6 +963,10 @@ def use_training_model_artifact(self) -> bool:
# otherwise, return true is a training model package is not set
return len(self.training_model_package_artifact_uris or {}) == 0

def is_gated_model(self) -> bool:
"""Returns True if the model has a EULA key or the model bucket is gated."""
return self.gated_bucket or self.hosting_eula_key is not None

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

hum, this is odd, shouldn't the pySDK only us gated_bucket as indicator?

Isn't it MH's unit test job to ensure that if hosting_eula_key is passed, then gated_bucket is set correctly?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We don't have a unit test for this (tho we should). It was confusing (at least for me) to refer to gated buckets and hosting eula keys separately, so to clear the confusion i created this helper function to clarify these are the markers for gated models (if the bucket is gated or there's a eula key).


def supports_incremental_training(self) -> bool:
"""Returns True if the model supports incremental training."""
return self.incremental_training_supported
Expand Down
22 changes: 13 additions & 9 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -476,21 +476,25 @@ def update_inference_tags_with_jumpstart_training_tags(
return inference_tags


def get_eula_message(model_specs: JumpStartModelSpecs, region: str) -> str:
"""Returns EULA message to display if one is available, else empty string."""
if model_specs.hosting_eula_key is None:
return ""
return (
f"Model '{model_specs.model_id}' requires accepting end-user license agreement (EULA). "
f"See https://{get_jumpstart_content_bucket(region=region)}.s3.{region}."
f"amazonaws.com{'.cn' if region.startswith('cn-') else ''}"
f"/{model_specs.hosting_eula_key} for terms of use."
)


def emit_logs_based_on_model_specs(
model_specs: JumpStartModelSpecs, region: str, s3_client: boto3.client
) -> None:
"""Emits logs based on model specs and region."""

if model_specs.hosting_eula_key:
constants.JUMPSTART_LOGGER.info(
"Model '%s' requires accepting end-user license agreement (EULA). "
"See https://%s.s3.%s.amazonaws.com%s/%s for terms of use.",
model_specs.model_id,
get_jumpstart_content_bucket(region=region),
region,
".cn" if region.startswith("cn-") else "",
model_specs.hosting_eula_key,
)
constants.JUMPSTART_LOGGER.info(get_eula_message(model_specs, region))

full_version: str = model_specs.version

Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -18,6 +18,7 @@
import pytest

from sagemaker import environment_variables
from sagemaker.jumpstart.utils import get_jumpstart_gated_content_bucket
from sagemaker.jumpstart.enums import JumpStartModelType

from tests.unit.sagemaker.jumpstart.utils import get_spec_from_base_spec, get_special_model_spec
Expand DownExpand Up@@ -203,6 +204,70 @@ def test_jumpstart_sdk_environment_variables(
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_1_artifact_all_variants(patched_get_model_specs):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model-1-artifact"
region = "us-west-2"

assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)


@patch("sagemaker.jumpstart.artifacts.environment_variables.JUMPSTART_LOGGER")
@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_no_gated_env_var_available(
patched_get_model_specs, patched_jumpstart_logger
):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model"
region = "us-west-2"

assert {} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)

patched_jumpstart_logger.warning.assert_called_once_with(
"'gemma-model' does not support ml.p3.2xlarge instance type for "
"training. Please use one of the following instance types: "
"ml.g5.12xlarge, ml.g5.24xlarge, ml.g5.48xlarge, ml.p4d.24xlarge."
)

# assert that supported instance types succeed
assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/g5/v1.0.0/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.g5.24xlarge",
script="training",
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_instance_type_overrides(patched_get_model_specs):

Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all
 blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks");
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Skip to content
34 changes: 32 additions & 2 deletions src/sagemaker/jumpstart/artifacts/environment_variables.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,10 +12,11 @@
# language governing permissions and limitations under the License.
"""This module contains functions for obtaining JumpStart environment variables."""
from __future__ import absolute_import
from typing import Dict, Optional
from typing import Callable, Dict, Optional, Set
from sagemaker.jumpstart.constants import (
DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
JUMPSTART_DEFAULT_REGION_NAME,
JUMPSTART_LOGGER,
SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY,
)
from sagemaker.jumpstart.enums import (
Expand DownExpand Up@@ -110,7 +111,9 @@ def _retrieve_default_environment_variables(

default_environment_variables.update(instance_specific_environment_variables)

gated_model_env_var: Optional[str] = _retrieve_gated_model_uri_env_var_value(
retrieve_gated_env_var_for_instance_type: Callable[
[str], Optional[str]
] = lambda instance_type: _retrieve_gated_model_uri_env_var_value(
model_id=model_id,
model_version=model_version,
region=region,
Expand All@@ -120,6 +123,33 @@ def _retrieve_default_environment_variables(
instance_type=instance_type,
)

gated_model_env_var: Optional[str] = retrieve_gated_env_var_for_instance_type(
instance_type
)

if gated_model_env_var is None and model_specs.is_gated_model():

possible_env_vars: Set[str] = {
retrieve_gated_env_var_for_instance_type(instance_type)
for instance_type in model_specs.supported_training_instance_types
}

# If all officially supported instance types have the same underlying artifact,
# we can use this artifact with high confidence that it'll succeed with
# an arbitrary instance.
if len(possible_env_vars) == 1:
gated_model_env_var = list(possible_env_vars)[0]

# If this model does not have 1 artifact for all supported instance types,
# we cannot determine which artifact to use for an arbitrary instance.
else:
log_msg = (
f"'{model_id}' does not support {instance_type} instance type"
" for training. Please use one of the following instance types: "
f"{', '.join(model_specs.supported_training_instance_types)}."
)
JUMPSTART_LOGGER.warning(log_msg)

if gated_model_env_var is not None:
default_environment_variables.update(
{SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY: gated_model_env_var}
Expand Down
21 changes: 21 additions & 0 deletions src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -62,6 +62,7 @@
)
from sagemaker.jumpstart.utils import (
add_jumpstart_model_id_version_tags,
get_eula_message,
update_dict_if_key_not_present,
resolve_estimator_sagemaker_config_field,
verify_model_region_and_return_specs,
Expand DownExpand Up@@ -597,6 +598,26 @@ def _add_env_to_kwargs(
value,
)

environment = getattr(kwargs, "environment", {}) or {}
if (
environment.get(SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY)
and str(environment.get("accept_eula", "")).lower() != "true"
):
model_specs = verify_model_region_and_return_specs(
model_id=kwargs.model_id,
version=kwargs.model_version,
region=kwargs.region,
scope=JumpStartScriptScope.TRAINING,
tolerate_deprecated_model=kwargs.tolerate_deprecated_model,
tolerate_vulnerable_model=kwargs.tolerate_vulnerable_model,
sagemaker_session=kwargs.sagemaker_session,
)
if model_specs.is_gated_model():

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Question: shouldn't this be the outer if statement here? And we would throw an error in all of the cases (the environment does not contain the special key, accept_eula is missing, accept_eula is false, etc.)?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'd rather make the conditions as specific as possible. Maybe we'll change how we launch gated models in the future.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

+1, this is odd, why bother retrieving the specs if you're not going to use them?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

fwiw, this won't involve another s3 call, it's just reading from memory

raise ValueError(
"Need to define ‘accept_eula'='true' within Environment. "
f"{get_eula_message(model_specs, kwargs.region)}"
)

return kwargs


Expand Down
4 changes: 4 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -963,6 +963,10 @@ def use_training_model_artifact(self) -> bool:
# otherwise, return true is a training model package is not set
return len(self.training_model_package_artifact_uris or {}) == 0

def is_gated_model(self) -> bool:
"""Returns True if the model has a EULA key or the model bucket is gated."""
return self.gated_bucket or self.hosting_eula_key is not None

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

hum, this is odd, shouldn't the pySDK only us gated_bucket as indicator?

Isn't it MH's unit test job to ensure that if hosting_eula_key is passed, then gated_bucket is set correctly?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We don't have a unit test for this (tho we should). It was confusing (at least for me) to refer to gated buckets and hosting eula keys separately, so to clear the confusion i created this helper function to clarify these are the markers for gated models (if the bucket is gated or there's a eula key).


def supports_incremental_training(self) -> bool:
"""Returns True if the model supports incremental training."""
return self.incremental_training_supported
Expand Down
22 changes: 13 additions & 9 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -476,21 +476,25 @@ def update_inference_tags_with_jumpstart_training_tags(
return inference_tags


def get_eula_message(model_specs: JumpStartModelSpecs, region: str) -> str:
"""Returns EULA message to display if one is available, else empty string."""
if model_specs.hosting_eula_key is None:
return ""
return (
f"Model '{model_specs.model_id}' requires accepting end-user license agreement (EULA). "
f"See https://{get_jumpstart_content_bucket(region=region)}.s3.{region}."
f"amazonaws.com{'.cn' if region.startswith('cn-') else ''}"
f"/{model_specs.hosting_eula_key} for terms of use."
)


def emit_logs_based_on_model_specs(
model_specs: JumpStartModelSpecs, region: str, s3_client: boto3.client
) -> None:
"""Emits logs based on model specs and region."""

if model_specs.hosting_eula_key:
constants.JUMPSTART_LOGGER.info(
"Model '%s' requires accepting end-user license agreement (EULA). "
"See https://%s.s3.%s.amazonaws.com%s/%s for terms of use.",
model_specs.model_id,
get_jumpstart_content_bucket(region=region),
region,
".cn" if region.startswith("cn-") else "",
model_specs.hosting_eula_key,
)
constants.JUMPSTART_LOGGER.info(get_eula_message(model_specs, region))

full_version: str = model_specs.version

Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -18,6 +18,7 @@
import pytest

from sagemaker import environment_variables
from sagemaker.jumpstart.utils import get_jumpstart_gated_content_bucket
from sagemaker.jumpstart.enums import JumpStartModelType

from tests.unit.sagemaker.jumpstart.utils import get_spec_from_base_spec, get_special_model_spec
Expand DownExpand Up@@ -203,6 +204,70 @@ def test_jumpstart_sdk_environment_variables(
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_1_artifact_all_variants(patched_get_model_specs):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model-1-artifact"
region = "us-west-2"

assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)


@patch("sagemaker.jumpstart.artifacts.environment_variables.JUMPSTART_LOGGER")
@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_no_gated_env_var_available(
patched_get_model_specs, patched_jumpstart_logger
):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model"
region = "us-west-2"

assert {} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)

patched_jumpstart_logger.warning.assert_called_once_with(
"'gemma-model' does not support ml.p3.2xlarge instance type for "
"training. Please use one of the following instance types: "
"ml.g5.12xlarge, ml.g5.24xlarge, ml.g5.48xlarge, ml.p4d.24xlarge."
)

# assert that supported instance types succeed
assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/g5/v1.0.0/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.g5.24xlarge",
script="training",
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_instance_type_overrides(patched_get_model_specs):

Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
34 changes: 32 additions & 2 deletions src/sagemaker/jumpstart/artifacts/environment_variables.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,10 +12,11 @@
# language governing permissions and limitations under the License.
"""This module contains functions for obtaining JumpStart environment variables."""
from __future__ import absolute_import
from typing import Dict, Optional
from typing import Callable, Dict, Optional, Set
from sagemaker.jumpstart.constants import (
DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
JUMPSTART_DEFAULT_REGION_NAME,
JUMPSTART_LOGGER,
SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY,
)
from sagemaker.jumpstart.enums import (
Expand DownExpand Up@@ -110,7 +111,9 @@ def _retrieve_default_environment_variables(

default_environment_variables.update(instance_specific_environment_variables)

gated_model_env_var: Optional[str] = _retrieve_gated_model_uri_env_var_value(
retrieve_gated_env_var_for_instance_type: Callable[
[str], Optional[str]
] = lambda instance_type: _retrieve_gated_model_uri_env_var_value(
model_id=model_id,
model_version=model_version,
region=region,
Expand All@@ -120,6 +123,33 @@ def _retrieve_default_environment_variables(
instance_type=instance_type,
)

gated_model_env_var: Optional[str] = retrieve_gated_env_var_for_instance_type(
instance_type
)

if gated_model_env_var is None and model_specs.is_gated_model():

possible_env_vars: Set[str] = {
retrieve_gated_env_var_for_instance_type(instance_type)
for instance_type in model_specs.supported_training_instance_types
}

# If all officially supported instance types have the same underlying artifact,
# we can use this artifact with high confidence that it'll succeed with
# an arbitrary instance.
if len(possible_env_vars) == 1:
gated_model_env_var = list(possible_env_vars)[0]

# If this model does not have 1 artifact for all supported instance types,
# we cannot determine which artifact to use for an arbitrary instance.
else:
log_msg = (
f"'{model_id}' does not support {instance_type} instance type"
" for training. Please use one of the following instance types: "
f"{', '.join(model_specs.supported_training_instance_types)}."
)
JUMPSTART_LOGGER.warning(log_msg)

if gated_model_env_var is not None:
default_environment_variables.update(
{SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY: gated_model_env_var}
Expand Down
21 changes: 21 additions & 0 deletions src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -62,6 +62,7 @@
)
from sagemaker.jumpstart.utils import (
add_jumpstart_model_id_version_tags,
get_eula_message,
update_dict_if_key_not_present,
resolve_estimator_sagemaker_config_field,
verify_model_region_and_return_specs,
Expand DownExpand Up@@ -597,6 +598,26 @@ def _add_env_to_kwargs(
value,
)

environment = getattr(kwargs, "environment", {}) or {}
if (
environment.get(SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY)
and str(environment.get("accept_eula", "")).lower() != "true"
):
model_specs = verify_model_region_and_return_specs(
model_id=kwargs.model_id,
version=kwargs.model_version,
region=kwargs.region,
scope=JumpStartScriptScope.TRAINING,
tolerate_deprecated_model=kwargs.tolerate_deprecated_model,
tolerate_vulnerable_model=kwargs.tolerate_vulnerable_model,
sagemaker_session=kwargs.sagemaker_session,
)
if model_specs.is_gated_model():

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Question: shouldn't this be the outer if statement here? And we would throw an error in all of the cases (the environment does not contain the special key, accept_eula is missing, accept_eula is false, etc.)?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'd rather make the conditions as specific as possible. Maybe we'll change how we launch gated models in the future.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

+1, this is odd, why bother retrieving the specs if you're not going to use them?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

fwiw, this won't involve another s3 call, it's just reading from memory

raise ValueError(
"Need to define ‘accept_eula'='true' within Environment. "
f"{get_eula_message(model_specs, kwargs.region)}"
)

return kwargs


Expand Down
4 changes: 4 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -963,6 +963,10 @@ def use_training_model_artifact(self) -> bool:
# otherwise, return true is a training model package is not set
return len(self.training_model_package_artifact_uris or {}) == 0

def is_gated_model(self) -> bool:
"""Returns True if the model has a EULA key or the model bucket is gated."""
return self.gated_bucket or self.hosting_eula_key is not None

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

hum, this is odd, shouldn't the pySDK only us gated_bucket as indicator?

Isn't it MH's unit test job to ensure that if hosting_eula_key is passed, then gated_bucket is set correctly?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We don't have a unit test for this (tho we should). It was confusing (at least for me) to refer to gated buckets and hosting eula keys separately, so to clear the confusion i created this helper function to clarify these are the markers for gated models (if the bucket is gated or there's a eula key).


def supports_incremental_training(self) -> bool:
"""Returns True if the model supports incremental training."""
return self.incremental_training_supported
Expand Down
22 changes: 13 additions & 9 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -476,21 +476,25 @@ def update_inference_tags_with_jumpstart_training_tags(
return inference_tags


def get_eula_message(model_specs: JumpStartModelSpecs, region: str) -> str:
"""Returns EULA message to display if one is available, else empty string."""
if model_specs.hosting_eula_key is None:
return ""
return (
f"Model '{model_specs.model_id}' requires accepting end-user license agreement (EULA). "
f"See https://{get_jumpstart_content_bucket(region=region)}.s3.{region}."
f"amazonaws.com{'.cn' if region.startswith('cn-') else ''}"
f"/{model_specs.hosting_eula_key} for terms of use."
)


def emit_logs_based_on_model_specs(
model_specs: JumpStartModelSpecs, region: str, s3_client: boto3.client
) -> None:
"""Emits logs based on model specs and region."""

if model_specs.hosting_eula_key:
constants.JUMPSTART_LOGGER.info(
"Model '%s' requires accepting end-user license agreement (EULA). "
"See https://%s.s3.%s.amazonaws.com%s/%s for terms of use.",
model_specs.model_id,
get_jumpstart_content_bucket(region=region),
region,
".cn" if region.startswith("cn-") else "",
model_specs.hosting_eula_key,
)
constants.JUMPSTART_LOGGER.info(get_eula_message(model_specs, region))

full_version: str = model_specs.version

Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -18,6 +18,7 @@
import pytest

from sagemaker import environment_variables
from sagemaker.jumpstart.utils import get_jumpstart_gated_content_bucket
from sagemaker.jumpstart.enums import JumpStartModelType

from tests.unit.sagemaker.jumpstart.utils import get_spec_from_base_spec, get_special_model_spec
Expand DownExpand Up@@ -203,6 +204,70 @@ def test_jumpstart_sdk_environment_variables(
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_1_artifact_all_variants(patched_get_model_specs):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model-1-artifact"
region = "us-west-2"

assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)


@patch("sagemaker.jumpstart.artifacts.environment_variables.JUMPSTART_LOGGER")
@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_no_gated_env_var_available(
patched_get_model_specs, patched_jumpstart_logger
):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model"
region = "us-west-2"

assert {} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)

patched_jumpstart_logger.warning.assert_called_once_with(
"'gemma-model' does not support ml.p3.2xlarge instance type for "
"training. Please use one of the following instance types: "
"ml.g5.12xlarge, ml.g5.24xlarge, ml.g5.48xlarge, ml.p4d.24xlarge."
)

# assert that supported instance types succeed
assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/g5/v1.0.0/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.g5.24xlarge",
script="training",
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_instance_type_overrides(patched_get_model_specs):

Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length > 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
34 changes: 32 additions & 2 deletions src/sagemaker/jumpstart/artifacts/environment_variables.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,10 +12,11 @@
# language governing permissions and limitations under the License.
"""This module contains functions for obtaining JumpStart environment variables."""
from __future__ import absolute_import
from typing import Dict, Optional
from typing import Callable, Dict, Optional, Set
from sagemaker.jumpstart.constants import (
DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
JUMPSTART_DEFAULT_REGION_NAME,
JUMPSTART_LOGGER,
SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY,
)
from sagemaker.jumpstart.enums import (
Expand DownExpand Up@@ -110,7 +111,9 @@ def _retrieve_default_environment_variables(

default_environment_variables.update(instance_specific_environment_variables)

gated_model_env_var: Optional[str] = _retrieve_gated_model_uri_env_var_value(
retrieve_gated_env_var_for_instance_type: Callable[
[str], Optional[str]
] = lambda instance_type: _retrieve_gated_model_uri_env_var_value(
model_id=model_id,
model_version=model_version,
region=region,
Expand All@@ -120,6 +123,33 @@ def _retrieve_default_environment_variables(
instance_type=instance_type,
)

gated_model_env_var: Optional[str] = retrieve_gated_env_var_for_instance_type(
instance_type
)

if gated_model_env_var is None and model_specs.is_gated_model():

possible_env_vars: Set[str] = {
retrieve_gated_env_var_for_instance_type(instance_type)
for instance_type in model_specs.supported_training_instance_types
}

# If all officially supported instance types have the same underlying artifact,
# we can use this artifact with high confidence that it'll succeed with
# an arbitrary instance.
if len(possible_env_vars) == 1:
gated_model_env_var = list(possible_env_vars)[0]

# If this model does not have 1 artifact for all supported instance types,
# we cannot determine which artifact to use for an arbitrary instance.
else:
log_msg = (
f"'{model_id}' does not support {instance_type} instance type"
" for training. Please use one of the following instance types: "
f"{', '.join(model_specs.supported_training_instance_types)}."
)
JUMPSTART_LOGGER.warning(log_msg)

if gated_model_env_var is not None:
default_environment_variables.update(
{SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY: gated_model_env_var}
Expand Down
21 changes: 21 additions & 0 deletions src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -62,6 +62,7 @@
)
from sagemaker.jumpstart.utils import (
add_jumpstart_model_id_version_tags,
get_eula_message,
update_dict_if_key_not_present,
resolve_estimator_sagemaker_config_field,
verify_model_region_and_return_specs,
Expand DownExpand Up@@ -597,6 +598,26 @@ def _add_env_to_kwargs(
value,
)

environment = getattr(kwargs, "environment", {}) or {}
if (
environment.get(SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY)
and str(environment.get("accept_eula", "")).lower() != "true"
):
model_specs = verify_model_region_and_return_specs(
model_id=kwargs.model_id,
version=kwargs.model_version,
region=kwargs.region,
scope=JumpStartScriptScope.TRAINING,
tolerate_deprecated_model=kwargs.tolerate_deprecated_model,
tolerate_vulnerable_model=kwargs.tolerate_vulnerable_model,
sagemaker_session=kwargs.sagemaker_session,
)
if model_specs.is_gated_model():

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Question: shouldn't this be the outer if statement here? And we would throw an error in all of the cases (the environment does not contain the special key, accept_eula is missing, accept_eula is false, etc.)?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'd rather make the conditions as specific as possible. Maybe we'll change how we launch gated models in the future.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

+1, this is odd, why bother retrieving the specs if you're not going to use them?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

fwiw, this won't involve another s3 call, it's just reading from memory

raise ValueError(
"Need to define ‘accept_eula'='true' within Environment. "
f"{get_eula_message(model_specs, kwargs.region)}"
)

return kwargs


Expand Down
4 changes: 4 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -963,6 +963,10 @@ def use_training_model_artifact(self) -> bool:
# otherwise, return true is a training model package is not set
return len(self.training_model_package_artifact_uris or {}) == 0

def is_gated_model(self) -> bool:
"""Returns True if the model has a EULA key or the model bucket is gated."""
return self.gated_bucket or self.hosting_eula_key is not None

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

hum, this is odd, shouldn't the pySDK only us gated_bucket as indicator?

Isn't it MH's unit test job to ensure that if hosting_eula_key is passed, then gated_bucket is set correctly?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We don't have a unit test for this (tho we should). It was confusing (at least for me) to refer to gated buckets and hosting eula keys separately, so to clear the confusion i created this helper function to clarify these are the markers for gated models (if the bucket is gated or there's a eula key).


def supports_incremental_training(self) -> bool:
"""Returns True if the model supports incremental training."""
return self.incremental_training_supported
Expand Down
22 changes: 13 additions & 9 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -476,21 +476,25 @@ def update_inference_tags_with_jumpstart_training_tags(
return inference_tags


def get_eula_message(model_specs: JumpStartModelSpecs, region: str) -> str:
"""Returns EULA message to display if one is available, else empty string."""
if model_specs.hosting_eula_key is None:
return ""
return (
f"Model '{model_specs.model_id}' requires accepting end-user license agreement (EULA). "
f"See https://{get_jumpstart_content_bucket(region=region)}.s3.{region}."
f"amazonaws.com{'.cn' if region.startswith('cn-') else ''}"
f"/{model_specs.hosting_eula_key} for terms of use."
)


def emit_logs_based_on_model_specs(
model_specs: JumpStartModelSpecs, region: str, s3_client: boto3.client
) -> None:
"""Emits logs based on model specs and region."""

if model_specs.hosting_eula_key:
constants.JUMPSTART_LOGGER.info(
"Model '%s' requires accepting end-user license agreement (EULA). "
"See https://%s.s3.%s.amazonaws.com%s/%s for terms of use.",
model_specs.model_id,
get_jumpstart_content_bucket(region=region),
region,
".cn" if region.startswith("cn-") else "",
model_specs.hosting_eula_key,
)
constants.JUMPSTART_LOGGER.info(get_eula_message(model_specs, region))

full_version: str = model_specs.version

Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -18,6 +18,7 @@
import pytest

from sagemaker import environment_variables
from sagemaker.jumpstart.utils import get_jumpstart_gated_content_bucket
from sagemaker.jumpstart.enums import JumpStartModelType

from tests.unit.sagemaker.jumpstart.utils import get_spec_from_base_spec, get_special_model_spec
Expand DownExpand Up@@ -203,6 +204,70 @@ def test_jumpstart_sdk_environment_variables(
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_1_artifact_all_variants(patched_get_model_specs):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model-1-artifact"
region = "us-west-2"

assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)


@patch("sagemaker.jumpstart.artifacts.environment_variables.JUMPSTART_LOGGER")
@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_no_gated_env_var_available(
patched_get_model_specs, patched_jumpstart_logger
):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model"
region = "us-west-2"

assert {} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)

patched_jumpstart_logger.warning.assert_called_once_with(
"'gemma-model' does not support ml.p3.2xlarge instance type for "
"training. Please use one of the following instance types: "
"ml.g5.12xlarge, ml.g5.24xlarge, ml.g5.48xlarge, ml.p4d.24xlarge."
)

# assert that supported instance types succeed
assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/g5/v1.0.0/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.g5.24xlarge",
script="training",
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_instance_type_overrides(patched_get_model_specs):

Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
Skip to content
34 changes: 32 additions & 2 deletions src/sagemaker/jumpstart/artifacts/environment_variables.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,10 +12,11 @@
# language governing permissions and limitations under the License.
"""This module contains functions for obtaining JumpStart environment variables."""
from __future__ import absolute_import
from typing import Dict, Optional
from typing import Callable, Dict, Optional, Set
from sagemaker.jumpstart.constants import (
DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
JUMPSTART_DEFAULT_REGION_NAME,
JUMPSTART_LOGGER,
SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY,
)
from sagemaker.jumpstart.enums import (
Expand DownExpand Up@@ -110,7 +111,9 @@ def _retrieve_default_environment_variables(

default_environment_variables.update(instance_specific_environment_variables)

gated_model_env_var: Optional[str] = _retrieve_gated_model_uri_env_var_value(
retrieve_gated_env_var_for_instance_type: Callable[
[str], Optional[str]
] = lambda instance_type: _retrieve_gated_model_uri_env_var_value(
model_id=model_id,
model_version=model_version,
region=region,
Expand All@@ -120,6 +123,33 @@ def _retrieve_default_environment_variables(
instance_type=instance_type,
)

gated_model_env_var: Optional[str] = retrieve_gated_env_var_for_instance_type(
instance_type
)

if gated_model_env_var is None and model_specs.is_gated_model():

possible_env_vars: Set[str] = {
retrieve_gated_env_var_for_instance_type(instance_type)
for instance_type in model_specs.supported_training_instance_types
}

# If all officially supported instance types have the same underlying artifact,
# we can use this artifact with high confidence that it'll succeed with
# an arbitrary instance.
if len(possible_env_vars) == 1:
gated_model_env_var = list(possible_env_vars)[0]

# If this model does not have 1 artifact for all supported instance types,
# we cannot determine which artifact to use for an arbitrary instance.
else:
log_msg = (
f"'{model_id}' does not support {instance_type} instance type"
" for training. Please use one of the following instance types: "
f"{', '.join(model_specs.supported_training_instance_types)}."
)
JUMPSTART_LOGGER.warning(log_msg)

if gated_model_env_var is not None:
default_environment_variables.update(
{SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY: gated_model_env_var}
Expand Down
21 changes: 21 additions & 0 deletions src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -62,6 +62,7 @@
)
from sagemaker.jumpstart.utils import (
add_jumpstart_model_id_version_tags,
get_eula_message,
update_dict_if_key_not_present,
resolve_estimator_sagemaker_config_field,
verify_model_region_and_return_specs,
Expand DownExpand Up@@ -597,6 +598,26 @@ def _add_env_to_kwargs(
value,
)

environment = getattr(kwargs, "environment", {}) or {}
if (
environment.get(SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY)
and str(environment.get("accept_eula", "")).lower() != "true"
):
model_specs = verify_model_region_and_return_specs(
model_id=kwargs.model_id,
version=kwargs.model_version,
region=kwargs.region,
scope=JumpStartScriptScope.TRAINING,
tolerate_deprecated_model=kwargs.tolerate_deprecated_model,
tolerate_vulnerable_model=kwargs.tolerate_vulnerable_model,
sagemaker_session=kwargs.sagemaker_session,
)
if model_specs.is_gated_model():

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Question: shouldn't this be the outer if statement here? And we would throw an error in all of the cases (the environment does not contain the special key, accept_eula is missing, accept_eula is false, etc.)?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'd rather make the conditions as specific as possible. Maybe we'll change how we launch gated models in the future.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

+1, this is odd, why bother retrieving the specs if you're not going to use them?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

fwiw, this won't involve another s3 call, it's just reading from memory

raise ValueError(
"Need to define ‘accept_eula'='true' within Environment. "
f"{get_eula_message(model_specs, kwargs.region)}"
)

return kwargs


Expand Down
4 changes: 4 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -963,6 +963,10 @@ def use_training_model_artifact(self) -> bool:
# otherwise, return true is a training model package is not set
return len(self.training_model_package_artifact_uris or {}) == 0

def is_gated_model(self) -> bool:
"""Returns True if the model has a EULA key or the model bucket is gated."""
return self.gated_bucket or self.hosting_eula_key is not None

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

hum, this is odd, shouldn't the pySDK only us gated_bucket as indicator?

Isn't it MH's unit test job to ensure that if hosting_eula_key is passed, then gated_bucket is set correctly?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We don't have a unit test for this (tho we should). It was confusing (at least for me) to refer to gated buckets and hosting eula keys separately, so to clear the confusion i created this helper function to clarify these are the markers for gated models (if the bucket is gated or there's a eula key).


def supports_incremental_training(self) -> bool:
"""Returns True if the model supports incremental training."""
return self.incremental_training_supported
Expand Down
22 changes: 13 additions & 9 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -476,21 +476,25 @@ def update_inference_tags_with_jumpstart_training_tags(
return inference_tags


def get_eula_message(model_specs: JumpStartModelSpecs, region: str) -> str:
"""Returns EULA message to display if one is available, else empty string."""
if model_specs.hosting_eula_key is None:
return ""
return (
f"Model '{model_specs.model_id}' requires accepting end-user license agreement (EULA). "
f"See https://{get_jumpstart_content_bucket(region=region)}.s3.{region}."
f"amazonaws.com{'.cn' if region.startswith('cn-') else ''}"
f"/{model_specs.hosting_eula_key} for terms of use."
)


def emit_logs_based_on_model_specs(
model_specs: JumpStartModelSpecs, region: str, s3_client: boto3.client
) -> None:
"""Emits logs based on model specs and region."""

if model_specs.hosting_eula_key:
constants.JUMPSTART_LOGGER.info(
"Model '%s' requires accepting end-user license agreement (EULA). "
"See https://%s.s3.%s.amazonaws.com%s/%s for terms of use.",
model_specs.model_id,
get_jumpstart_content_bucket(region=region),
region,
".cn" if region.startswith("cn-") else "",
model_specs.hosting_eula_key,
)
constants.JUMPSTART_LOGGER.info(get_eula_message(model_specs, region))

full_version: str = model_specs.version

Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -18,6 +18,7 @@
import pytest

from sagemaker import environment_variables
from sagemaker.jumpstart.utils import get_jumpstart_gated_content_bucket
from sagemaker.jumpstart.enums import JumpStartModelType

from tests.unit.sagemaker.jumpstart.utils import get_spec_from_base_spec, get_special_model_spec
Expand DownExpand Up@@ -203,6 +204,70 @@ def test_jumpstart_sdk_environment_variables(
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_1_artifact_all_variants(patched_get_model_specs):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model-1-artifact"
region = "us-west-2"

assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)


@patch("sagemaker.jumpstart.artifacts.environment_variables.JUMPSTART_LOGGER")
@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_no_gated_env_var_available(
patched_get_model_specs, patched_jumpstart_logger
):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model"
region = "us-west-2"

assert {} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)

patched_jumpstart_logger.warning.assert_called_once_with(
"'gemma-model' does not support ml.p3.2xlarge instance type for "
"training. Please use one of the following instance types: "
"ml.g5.12xlarge, ml.g5.24xlarge, ml.g5.48xlarge, ml.p4d.24xlarge."
)

# assert that supported instance types succeed
assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/g5/v1.0.0/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.g5.24xlarge",
script="training",
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_instance_type_overrides(patched_get_model_specs):

Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
34 changes: 32 additions & 2 deletions src/sagemaker/jumpstart/artifacts/environment_variables.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,10 +12,11 @@
# language governing permissions and limitations under the License.
"""This module contains functions for obtaining JumpStart environment variables."""
from __future__ import absolute_import
from typing import Dict, Optional
from typing import Callable, Dict, Optional, Set
from sagemaker.jumpstart.constants import (
DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
JUMPSTART_DEFAULT_REGION_NAME,
JUMPSTART_LOGGER,
SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY,
)
from sagemaker.jumpstart.enums import (
Expand DownExpand Up@@ -110,7 +111,9 @@ def _retrieve_default_environment_variables(

default_environment_variables.update(instance_specific_environment_variables)

gated_model_env_var: Optional[str] = _retrieve_gated_model_uri_env_var_value(
retrieve_gated_env_var_for_instance_type: Callable[
[str], Optional[str]
] = lambda instance_type: _retrieve_gated_model_uri_env_var_value(
model_id=model_id,
model_version=model_version,
region=region,
Expand All@@ -120,6 +123,33 @@ def _retrieve_default_environment_variables(
instance_type=instance_type,
)

gated_model_env_var: Optional[str] = retrieve_gated_env_var_for_instance_type(
instance_type
)

if gated_model_env_var is None and model_specs.is_gated_model():

possible_env_vars: Set[str] = {
retrieve_gated_env_var_for_instance_type(instance_type)
for instance_type in model_specs.supported_training_instance_types
}

# If all officially supported instance types have the same underlying artifact,
# we can use this artifact with high confidence that it'll succeed with
# an arbitrary instance.
if len(possible_env_vars) == 1:
gated_model_env_var = list(possible_env_vars)[0]

# If this model does not have 1 artifact for all supported instance types,
# we cannot determine which artifact to use for an arbitrary instance.
else:
log_msg = (
f"'{model_id}' does not support {instance_type} instance type"
" for training. Please use one of the following instance types: "
f"{', '.join(model_specs.supported_training_instance_types)}."
)
JUMPSTART_LOGGER.warning(log_msg)

if gated_model_env_var is not None:
default_environment_variables.update(
{SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY: gated_model_env_var}
Expand Down
21 changes: 21 additions & 0 deletions src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -62,6 +62,7 @@
)
from sagemaker.jumpstart.utils import (
add_jumpstart_model_id_version_tags,
get_eula_message,
update_dict_if_key_not_present,
resolve_estimator_sagemaker_config_field,
verify_model_region_and_return_specs,
Expand DownExpand Up@@ -597,6 +598,26 @@ def _add_env_to_kwargs(
value,
)

environment = getattr(kwargs, "environment", {}) or {}
if (
environment.get(SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY)
and str(environment.get("accept_eula", "")).lower() != "true"
):
model_specs = verify_model_region_and_return_specs(
model_id=kwargs.model_id,
version=kwargs.model_version,
region=kwargs.region,
scope=JumpStartScriptScope.TRAINING,
tolerate_deprecated_model=kwargs.tolerate_deprecated_model,
tolerate_vulnerable_model=kwargs.tolerate_vulnerable_model,
sagemaker_session=kwargs.sagemaker_session,
)
if model_specs.is_gated_model():

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Question: shouldn't this be the outer if statement here? And we would throw an error in all of the cases (the environment does not contain the special key, accept_eula is missing, accept_eula is false, etc.)?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'd rather make the conditions as specific as possible. Maybe we'll change how we launch gated models in the future.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

+1, this is odd, why bother retrieving the specs if you're not going to use them?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

fwiw, this won't involve another s3 call, it's just reading from memory

raise ValueError(
"Need to define ‘accept_eula'='true' within Environment. "
f"{get_eula_message(model_specs, kwargs.region)}"
)

return kwargs


Expand Down
4 changes: 4 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -963,6 +963,10 @@ def use_training_model_artifact(self) -> bool:
# otherwise, return true is a training model package is not set
return len(self.training_model_package_artifact_uris or {}) == 0

def is_gated_model(self) -> bool:
"""Returns True if the model has a EULA key or the model bucket is gated."""
return self.gated_bucket or self.hosting_eula_key is not None

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

hum, this is odd, shouldn't the pySDK only us gated_bucket as indicator?

Isn't it MH's unit test job to ensure that if hosting_eula_key is passed, then gated_bucket is set correctly?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We don't have a unit test for this (tho we should). It was confusing (at least for me) to refer to gated buckets and hosting eula keys separately, so to clear the confusion i created this helper function to clarify these are the markers for gated models (if the bucket is gated or there's a eula key).


def supports_incremental_training(self) -> bool:
"""Returns True if the model supports incremental training."""
return self.incremental_training_supported
Expand Down
22 changes: 13 additions & 9 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -476,21 +476,25 @@ def update_inference_tags_with_jumpstart_training_tags(
return inference_tags


def get_eula_message(model_specs: JumpStartModelSpecs, region: str) -> str:
"""Returns EULA message to display if one is available, else empty string."""
if model_specs.hosting_eula_key is None:
return ""
return (
f"Model '{model_specs.model_id}' requires accepting end-user license agreement (EULA). "
f"See https://{get_jumpstart_content_bucket(region=region)}.s3.{region}."
f"amazonaws.com{'.cn' if region.startswith('cn-') else ''}"
f"/{model_specs.hosting_eula_key} for terms of use."
)


def emit_logs_based_on_model_specs(
model_specs: JumpStartModelSpecs, region: str, s3_client: boto3.client
) -> None:
"""Emits logs based on model specs and region."""

if model_specs.hosting_eula_key:
constants.JUMPSTART_LOGGER.info(
"Model '%s' requires accepting end-user license agreement (EULA). "
"See https://%s.s3.%s.amazonaws.com%s/%s for terms of use.",
model_specs.model_id,
get_jumpstart_content_bucket(region=region),
region,
".cn" if region.startswith("cn-") else "",
model_specs.hosting_eula_key,
)
constants.JUMPSTART_LOGGER.info(get_eula_message(model_specs, region))

full_version: str = model_specs.version

Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -18,6 +18,7 @@
import pytest

from sagemaker import environment_variables
from sagemaker.jumpstart.utils import get_jumpstart_gated_content_bucket
from sagemaker.jumpstart.enums import JumpStartModelType

from tests.unit.sagemaker.jumpstart.utils import get_spec_from_base_spec, get_special_model_spec
Expand DownExpand Up@@ -203,6 +204,70 @@ def test_jumpstart_sdk_environment_variables(
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_1_artifact_all_variants(patched_get_model_specs):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model-1-artifact"
region = "us-west-2"

assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)


@patch("sagemaker.jumpstart.artifacts.environment_variables.JUMPSTART_LOGGER")
@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_no_gated_env_var_available(
patched_get_model_specs, patched_jumpstart_logger
):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model"
region = "us-west-2"

assert {} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)

patched_jumpstart_logger.warning.assert_called_once_with(
"'gemma-model' does not support ml.p3.2xlarge instance type for "
"training. Please use one of the following instance types: "
"ml.g5.12xlarge, ml.g5.24xlarge, ml.g5.48xlarge, ml.p4d.24xlarge."
)

# assert that supported instance types succeed
assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/g5/v1.0.0/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.g5.24xlarge",
script="training",
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_instance_type_overrides(patched_get_model_specs):

Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
Skip to content
34 changes: 32 additions & 2 deletions src/sagemaker/jumpstart/artifacts/environment_variables.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,10 +12,11 @@
# language governing permissions and limitations under the License.
"""This module contains functions for obtaining JumpStart environment variables."""
from __future__ import absolute_import
from typing import Dict, Optional
from typing import Callable, Dict, Optional, Set
from sagemaker.jumpstart.constants import (
DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
JUMPSTART_DEFAULT_REGION_NAME,
JUMPSTART_LOGGER,
SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY,
)
from sagemaker.jumpstart.enums import (
Expand DownExpand Up@@ -110,7 +111,9 @@ def _retrieve_default_environment_variables(

default_environment_variables.update(instance_specific_environment_variables)

gated_model_env_var: Optional[str] = _retrieve_gated_model_uri_env_var_value(
retrieve_gated_env_var_for_instance_type: Callable[
[str], Optional[str]
] = lambda instance_type: _retrieve_gated_model_uri_env_var_value(
model_id=model_id,
model_version=model_version,
region=region,
Expand All@@ -120,6 +123,33 @@ def _retrieve_default_environment_variables(
instance_type=instance_type,
)

gated_model_env_var: Optional[str] = retrieve_gated_env_var_for_instance_type(
instance_type
)

if gated_model_env_var is None and model_specs.is_gated_model():

possible_env_vars: Set[str] = {
retrieve_gated_env_var_for_instance_type(instance_type)
for instance_type in model_specs.supported_training_instance_types
}

# If all officially supported instance types have the same underlying artifact,
# we can use this artifact with high confidence that it'll succeed with
# an arbitrary instance.
if len(possible_env_vars) == 1:
gated_model_env_var = list(possible_env_vars)[0]

# If this model does not have 1 artifact for all supported instance types,
# we cannot determine which artifact to use for an arbitrary instance.
else:
log_msg = (
f"'{model_id}' does not support {instance_type} instance type"
" for training. Please use one of the following instance types: "
f"{', '.join(model_specs.supported_training_instance_types)}."
)
JUMPSTART_LOGGER.warning(log_msg)

if gated_model_env_var is not None:
default_environment_variables.update(
{SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY: gated_model_env_var}
Expand Down
21 changes: 21 additions & 0 deletions src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -62,6 +62,7 @@
)
from sagemaker.jumpstart.utils import (
add_jumpstart_model_id_version_tags,
get_eula_message,
update_dict_if_key_not_present,
resolve_estimator_sagemaker_config_field,
verify_model_region_and_return_specs,
Expand DownExpand Up@@ -597,6 +598,26 @@ def _add_env_to_kwargs(
value,
)

environment = getattr(kwargs, "environment", {}) or {}
if (
environment.get(SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY)
and str(environment.get("accept_eula", "")).lower() != "true"
):
model_specs = verify_model_region_and_return_specs(
model_id=kwargs.model_id,
version=kwargs.model_version,
region=kwargs.region,
scope=JumpStartScriptScope.TRAINING,
tolerate_deprecated_model=kwargs.tolerate_deprecated_model,
tolerate_vulnerable_model=kwargs.tolerate_vulnerable_model,
sagemaker_session=kwargs.sagemaker_session,
)
if model_specs.is_gated_model():

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Question: shouldn't this be the outer if statement here? And we would throw an error in all of the cases (the environment does not contain the special key, accept_eula is missing, accept_eula is false, etc.)?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'd rather make the conditions as specific as possible. Maybe we'll change how we launch gated models in the future.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

+1, this is odd, why bother retrieving the specs if you're not going to use them?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

fwiw, this won't involve another s3 call, it's just reading from memory

raise ValueError(
"Need to define ‘accept_eula'='true' within Environment. "
f"{get_eula_message(model_specs, kwargs.region)}"
)

return kwargs


Expand Down
4 changes: 4 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -963,6 +963,10 @@ def use_training_model_artifact(self) -> bool:
# otherwise, return true is a training model package is not set
return len(self.training_model_package_artifact_uris or {}) == 0

def is_gated_model(self) -> bool:
"""Returns True if the model has a EULA key or the model bucket is gated."""
return self.gated_bucket or self.hosting_eula_key is not None

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

hum, this is odd, shouldn't the pySDK only us gated_bucket as indicator?

Isn't it MH's unit test job to ensure that if hosting_eula_key is passed, then gated_bucket is set correctly?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We don't have a unit test for this (tho we should). It was confusing (at least for me) to refer to gated buckets and hosting eula keys separately, so to clear the confusion i created this helper function to clarify these are the markers for gated models (if the bucket is gated or there's a eula key).


def supports_incremental_training(self) -> bool:
"""Returns True if the model supports incremental training."""
return self.incremental_training_supported
Expand Down
22 changes: 13 additions & 9 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -476,21 +476,25 @@ def update_inference_tags_with_jumpstart_training_tags(
return inference_tags


def get_eula_message(model_specs: JumpStartModelSpecs, region: str) -> str:
"""Returns EULA message to display if one is available, else empty string."""
if model_specs.hosting_eula_key is None:
return ""
return (
f"Model '{model_specs.model_id}' requires accepting end-user license agreement (EULA). "
f"See https://{get_jumpstart_content_bucket(region=region)}.s3.{region}."
f"amazonaws.com{'.cn' if region.startswith('cn-') else ''}"
f"/{model_specs.hosting_eula_key} for terms of use."
)


def emit_logs_based_on_model_specs(
model_specs: JumpStartModelSpecs, region: str, s3_client: boto3.client
) -> None:
"""Emits logs based on model specs and region."""

if model_specs.hosting_eula_key:
constants.JUMPSTART_LOGGER.info(
"Model '%s' requires accepting end-user license agreement (EULA). "
"See https://%s.s3.%s.amazonaws.com%s/%s for terms of use.",
model_specs.model_id,
get_jumpstart_content_bucket(region=region),
region,
".cn" if region.startswith("cn-") else "",
model_specs.hosting_eula_key,
)
constants.JUMPSTART_LOGGER.info(get_eula_message(model_specs, region))

full_version: str = model_specs.version

Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -18,6 +18,7 @@
import pytest

from sagemaker import environment_variables
from sagemaker.jumpstart.utils import get_jumpstart_gated_content_bucket
from sagemaker.jumpstart.enums import JumpStartModelType

from tests.unit.sagemaker.jumpstart.utils import get_spec_from_base_spec, get_special_model_spec
Expand DownExpand Up@@ -203,6 +204,70 @@ def test_jumpstart_sdk_environment_variables(
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_1_artifact_all_variants(patched_get_model_specs):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model-1-artifact"
region = "us-west-2"

assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)


@patch("sagemaker.jumpstart.artifacts.environment_variables.JUMPSTART_LOGGER")
@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_no_gated_env_var_available(
patched_get_model_specs, patched_jumpstart_logger
):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model"
region = "us-west-2"

assert {} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)

patched_jumpstart_logger.warning.assert_called_once_with(
"'gemma-model' does not support ml.p3.2xlarge instance type for "
"training. Please use one of the following instance types: "
"ml.g5.12xlarge, ml.g5.24xlarge, ml.g5.48xlarge, ml.p4d.24xlarge."
)

# assert that supported instance types succeed
assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/g5/v1.0.0/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.g5.24xlarge",
script="training",
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_instance_type_overrides(patched_get_model_specs):

Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
Skip to content
34 changes: 32 additions & 2 deletions src/sagemaker/jumpstart/artifacts/environment_variables.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -12,10 +12,11 @@
# language governing permissions and limitations under the License.
"""This module contains functions for obtaining JumpStart environment variables."""
from __future__ import absolute_import
from typing import Dict, Optional
from typing import Callable, Dict, Optional, Set
from sagemaker.jumpstart.constants import (
DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
JUMPSTART_DEFAULT_REGION_NAME,
JUMPSTART_LOGGER,
SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY,
)
from sagemaker.jumpstart.enums import (
Expand DownExpand Up@@ -110,7 +111,9 @@ def _retrieve_default_environment_variables(

default_environment_variables.update(instance_specific_environment_variables)

gated_model_env_var: Optional[str] = _retrieve_gated_model_uri_env_var_value(
retrieve_gated_env_var_for_instance_type: Callable[
[str], Optional[str]
] = lambda instance_type: _retrieve_gated_model_uri_env_var_value(
model_id=model_id,
model_version=model_version,
region=region,
Expand All@@ -120,6 +123,33 @@ def _retrieve_default_environment_variables(
instance_type=instance_type,
)

gated_model_env_var: Optional[str] = retrieve_gated_env_var_for_instance_type(
instance_type
)

if gated_model_env_var is None and model_specs.is_gated_model():

possible_env_vars: Set[str] = {
retrieve_gated_env_var_for_instance_type(instance_type)
for instance_type in model_specs.supported_training_instance_types
}

# If all officially supported instance types have the same underlying artifact,
# we can use this artifact with high confidence that it'll succeed with
# an arbitrary instance.
if len(possible_env_vars) == 1:
gated_model_env_var = list(possible_env_vars)[0]

# If this model does not have 1 artifact for all supported instance types,
# we cannot determine which artifact to use for an arbitrary instance.
else:
log_msg = (
f"'{model_id}' does not support {instance_type} instance type"
" for training. Please use one of the following instance types: "
f"{', '.join(model_specs.supported_training_instance_types)}."
)
JUMPSTART_LOGGER.warning(log_msg)

if gated_model_env_var is not None:
default_environment_variables.update(
{SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY: gated_model_env_var}
Expand Down
21 changes: 21 additions & 0 deletions src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -62,6 +62,7 @@
)
from sagemaker.jumpstart.utils import (
add_jumpstart_model_id_version_tags,
get_eula_message,
update_dict_if_key_not_present,
resolve_estimator_sagemaker_config_field,
verify_model_region_and_return_specs,
Expand DownExpand Up@@ -597,6 +598,26 @@ def _add_env_to_kwargs(
value,
)

environment = getattr(kwargs, "environment", {}) or {}
if (
environment.get(SAGEMAKER_GATED_MODEL_S3_URI_TRAINING_ENV_VAR_KEY)
and str(environment.get("accept_eula", "")).lower() != "true"
):
model_specs = verify_model_region_and_return_specs(
model_id=kwargs.model_id,
version=kwargs.model_version,
region=kwargs.region,
scope=JumpStartScriptScope.TRAINING,
tolerate_deprecated_model=kwargs.tolerate_deprecated_model,
tolerate_vulnerable_model=kwargs.tolerate_vulnerable_model,
sagemaker_session=kwargs.sagemaker_session,
)
if model_specs.is_gated_model():

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Question: shouldn't this be the outer if statement here? And we would throw an error in all of the cases (the environment does not contain the special key, accept_eula is missing, accept_eula is false, etc.)?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'd rather make the conditions as specific as possible. Maybe we'll change how we launch gated models in the future.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

+1, this is odd, why bother retrieving the specs if you're not going to use them?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

fwiw, this won't involve another s3 call, it's just reading from memory

raise ValueError(
"Need to define ‘accept_eula'='true' within Environment. "
f"{get_eula_message(model_specs, kwargs.region)}"
)

return kwargs


Expand Down
4 changes: 4 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -963,6 +963,10 @@ def use_training_model_artifact(self) -> bool:
# otherwise, return true is a training model package is not set
return len(self.training_model_package_artifact_uris or {}) == 0

def is_gated_model(self) -> bool:
"""Returns True if the model has a EULA key or the model bucket is gated."""
return self.gated_bucket or self.hosting_eula_key is not None

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

hum, this is odd, shouldn't the pySDK only us gated_bucket as indicator?

Isn't it MH's unit test job to ensure that if hosting_eula_key is passed, then gated_bucket is set correctly?

Copy link
Copy Markdown
MemberAuthor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We don't have a unit test for this (tho we should). It was confusing (at least for me) to refer to gated buckets and hosting eula keys separately, so to clear the confusion i created this helper function to clarify these are the markers for gated models (if the bucket is gated or there's a eula key).


def supports_incremental_training(self) -> bool:
"""Returns True if the model supports incremental training."""
return self.incremental_training_supported
Expand Down
22 changes: 13 additions & 9 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -476,21 +476,25 @@ def update_inference_tags_with_jumpstart_training_tags(
return inference_tags


def get_eula_message(model_specs: JumpStartModelSpecs, region: str) -> str:
"""Returns EULA message to display if one is available, else empty string."""
if model_specs.hosting_eula_key is None:
return ""
return (
f"Model '{model_specs.model_id}' requires accepting end-user license agreement (EULA). "
f"See https://{get_jumpstart_content_bucket(region=region)}.s3.{region}."
f"amazonaws.com{'.cn' if region.startswith('cn-') else ''}"
f"/{model_specs.hosting_eula_key} for terms of use."
)


def emit_logs_based_on_model_specs(
model_specs: JumpStartModelSpecs, region: str, s3_client: boto3.client
) -> None:
"""Emits logs based on model specs and region."""

if model_specs.hosting_eula_key:
constants.JUMPSTART_LOGGER.info(
"Model '%s' requires accepting end-user license agreement (EULA). "
"See https://%s.s3.%s.amazonaws.com%s/%s for terms of use.",
model_specs.model_id,
get_jumpstart_content_bucket(region=region),
region,
".cn" if region.startswith("cn-") else "",
model_specs.hosting_eula_key,
)
constants.JUMPSTART_LOGGER.info(get_eula_message(model_specs, region))

full_version: str = model_specs.version

Expand Down
Original file line numberDiff line numberDiff line change
Expand Up@@ -18,6 +18,7 @@
import pytest

from sagemaker import environment_variables
from sagemaker.jumpstart.utils import get_jumpstart_gated_content_bucket
from sagemaker.jumpstart.enums import JumpStartModelType

from tests.unit.sagemaker.jumpstart.utils import get_spec_from_base_spec, get_special_model_spec
Expand DownExpand Up@@ -203,6 +204,70 @@ def test_jumpstart_sdk_environment_variables(
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_1_artifact_all_variants(patched_get_model_specs):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model-1-artifact"
region = "us-west-2"

assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)


@patch("sagemaker.jumpstart.artifacts.environment_variables.JUMPSTART_LOGGER")
@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_no_gated_env_var_available(
patched_get_model_specs, patched_jumpstart_logger
):

patched_get_model_specs.side_effect = get_special_model_spec

model_id = "gemma-model"
region = "us-west-2"

assert {} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.p3.2xlarge",
script="training",
)

patched_jumpstart_logger.warning.assert_called_once_with(
"'gemma-model' does not support ml.p3.2xlarge instance type for "
"training. Please use one of the following instance types: "
"ml.g5.12xlarge, ml.g5.24xlarge, ml.g5.48xlarge, ml.p4d.24xlarge."
)

# assert that supported instance types succeed
assert {
"SageMakerGatedModelS3Uri": f"s3://{get_jumpstart_gated_content_bucket(region)}/"
"huggingface-training/g5/v1.0.0/train-huggingface-llm-gemma-7b-instruct.tar.gz"
} == environment_variables.retrieve_default(
region=region,
model_id=model_id,
model_version="*",
include_aws_sdk_env_vars=False,
sagemaker_session=mock_session,
instance_type="ml.g5.24xlarge",
script="training",
)


@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_jumpstart_sdk_environment_variables_instance_type_overrides(patched_get_model_specs):

Expand Down
Loading