Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions src/sagemaker/jumpstart/estimator.py
Original file line number Diff line number Diff line change
Expand Up @@ -1206,6 +1206,8 @@ def list_training_configs(self) -> List[JumpStartMetadataConfig]:
region=self.region,
scope=JumpStartScriptScope.TRAINING,
sagemaker_session=self.sagemaker_session,
tolerate_vulnerable_model=self.tolerate_vulnerable_model,
tolerate_deprecated_model=self.tolerate_deprecated_model,
)
return list(configs_dict.values())

Expand Down
2 changes: 2 additions & 0 deletions src/sagemaker/jumpstart/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -401,6 +401,8 @@ def _validate_model_id_and_type():
sagemaker_session=self.sagemaker_session,
model_type=self.model_type,
hub_arn=self.hub_arn,
tolerate_vulnerable_model=self.tolerate_vulnerable_model,
tolerate_deprecated_model=self.tolerate_deprecated_model,
)

def log_subscription_warning(self) -> None:
Expand Down
13 changes: 13 additions & 0 deletions src/sagemaker/jumpstart/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1146,9 +1146,20 @@ def get_jumpstart_configs(
scope: enums.JumpStartScriptScope = enums.JumpStartScriptScope.INFERENCE,
model_type: enums.JumpStartModelType = enums.JumpStartModelType.OPEN_WEIGHTS,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
) -> Dict[str, JumpStartMetadataConfig]:
"""Returns metadata configs for the given model ID and region.

Args:
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
security vulnerabilities. (Default: False).
tolerate_deprecated_model (bool): True if deprecated models should be tolerated
(exception not raised). False if these models should raise an exception.
(Default: False).

Raises:
ValueError: If the script scope is not supported by JumpStart.
"""
Expand All @@ -1160,6 +1171,8 @@ def get_jumpstart_configs(
scope=scope,
model_type=model_type,
hub_arn=hub_arn,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
)

if scope == enums.JumpStartScriptScope.INFERENCE:
Expand Down
31 changes: 31 additions & 0 deletions tests/unit/sagemaker/jumpstart/estimator/test_estimator.py
Original file line number Diff line number Diff line change
Expand Up @@ -2301,6 +2301,37 @@ def test_estimator_set_config_name(

mock_estimator_fit.assert_called_once_with(wait=True, job_name="pt-ic-mobilenet-v2-8675309")

@mock.patch("sagemaker.jumpstart.estimator.get_jumpstart_configs")
@mock.patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor._get_manifest")
@mock.patch("sagemaker.jumpstart.factory.estimator.Session")
@mock.patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
@mock.patch("sagemaker.jumpstart.estimator.Estimator.__init__")
@mock.patch("sagemaker.jumpstart.factory.estimator.JUMPSTART_DEFAULT_REGION_NAME", region)
def test_estimator_list_training_configs_forwards_tolerance(
self,
mock_estimator_init: mock.Mock,
mock_get_model_specs: mock.Mock,
mock_session: mock.Mock,
mock_get_manifest: mock.Mock,
mock_get_jumpstart_configs: mock.Mock,
):
mock_get_model_specs.side_effect = get_prototype_spec_with_configs
mock_get_manifest.side_effect = (
lambda region, model_type, *args, **kwargs: get_prototype_manifest(region, model_type)
)
mock_get_jumpstart_configs.return_value = {}
mock_session.return_value = sagemaker_session

estimator = JumpStartEstimator(
model_id="pytorch-ic-mobilenet-v2",
tolerate_vulnerable_model=True,
tolerate_deprecated_model=True,
)

assert estimator.list_training_configs() == []
assert mock_get_jumpstart_configs.call_args.kwargs["tolerate_vulnerable_model"] is True
assert mock_get_jumpstart_configs.call_args.kwargs["tolerate_deprecated_model"] is True

@mock.patch("sagemaker.utils.sagemaker_timestamp")
@mock.patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor._get_manifest")
@mock.patch("sagemaker.jumpstart.factory.estimator.Session")
Expand Down
104 changes: 104 additions & 0 deletions tests/unit/sagemaker/jumpstart/test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1812,6 +1812,110 @@ def test_get_jumpstart_configs_success(
"ml.p3.2xlarge",
]

@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_get_jumpstart_configs_vulnerable_model_raises_by_default(
self,
patched_get_model_specs,
):
def make_vulnerable_spec(*largs, **kwargs):
spec = get_base_spec_with_prototype_configs()
spec.inference_vulnerable = True
spec.inference_vulnerabilities = ["CVE-2024-11393"]
return spec

patched_get_model_specs.side_effect = make_vulnerable_spec

with pytest.raises(VulnerableJumpStartModelError):
utils.get_jumpstart_configs(
"mock-region", "mock-model", "mock-model-version", config_names=["gpu-inference"]
)

@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_get_jumpstart_configs_tolerates_vulnerable_model(
self,
patched_get_model_specs,
):
def make_vulnerable_spec(*largs, **kwargs):
spec = get_base_spec_with_prototype_configs()
spec.inference_vulnerable = True
spec.inference_vulnerabilities = ["CVE-2024-11393"]
return spec

patched_get_model_specs.side_effect = make_vulnerable_spec

configs = utils.get_jumpstart_configs(
"mock-region",
"mock-model",
"mock-model-version",
config_names=["gpu-inference"],
tolerate_vulnerable_model=True,
)

assert configs.keys() == {"gpu-inference"}

@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_get_jumpstart_configs_deprecated_model_raises_by_default(
self,
patched_get_model_specs,
):
def make_deprecated_spec(*largs, **kwargs):
spec = get_base_spec_with_prototype_configs()
spec.deprecated = True
return spec

patched_get_model_specs.side_effect = make_deprecated_spec

with pytest.raises(DeprecatedJumpStartModelError):
utils.get_jumpstart_configs(
"mock-region", "mock-model", "mock-model-version", config_names=["gpu-inference"]
)

@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_get_jumpstart_configs_tolerates_deprecated_model(
self,
patched_get_model_specs,
):
def make_deprecated_spec(*largs, **kwargs):
spec = get_base_spec_with_prototype_configs()
spec.deprecated = True
return spec

patched_get_model_specs.side_effect = make_deprecated_spec

configs = utils.get_jumpstart_configs(
"mock-region",
"mock-model",
"mock-model-version",
config_names=["gpu-inference"],
tolerate_deprecated_model=True,
)

assert configs.keys() == {"gpu-inference"}

@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
def test_get_jumpstart_configs_tolerates_vulnerable_training_model(
self,
patched_get_model_specs,
):
def make_vulnerable_spec(*largs, **kwargs):
spec = get_base_spec_with_prototype_configs()
spec.training_vulnerable = True
spec.training_vulnerabilities = ["CVE-2024-11393"]
return spec

patched_get_model_specs.side_effect = make_vulnerable_spec

configs = utils.get_jumpstart_configs(
"mock-region",
"mock-model",
"mock-model-version",
config_names=["gpu-training"],
scope=JumpStartScriptScope.TRAINING,
tolerate_vulnerable_model=True,
)

assert configs.keys() == {"gpu-training"}


class TestBenchmarkStats:
@patch("sagemaker.jumpstart.accessors.JumpStartModelsAccessor.get_model_specs")
Expand Down
Loading