Merged
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
1 change: 1 addition & 0 deletions src/sagemaker/jumpstart/enums.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -92,6 +92,7 @@ class JumpStartTag(str, Enum):
MODEL_ID = "sagemaker-sdk:jumpstart-model-id"
MODEL_VERSION = "sagemaker-sdk:jumpstart-model-version"
MODEL_TYPE = "sagemaker-sdk:jumpstart-model-type"
MODEL_CONFIG_NAME = "sagemaker-sdk:jumpstart-model-config-name"


class SerializerType(str, Enum):
Expand Down
9 changes: 5 additions & 4 deletions src/sagemaker/jumpstart/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -33,7 +33,7 @@

from sagemaker.jumpstart.factory.estimator import get_deploy_kwargs, get_fit_kwargs, get_init_kwargs
from sagemaker.jumpstart.factory.model import get_default_predictor
from sagemaker.jumpstart.session_utils import get_model_id_version_from_training_job
from sagemaker.jumpstart.session_utils import get_model_info_from_training_job
from sagemaker.jumpstart.types import JumpStartMetadataConfig
from sagemaker.jumpstart.utils import (
get_jumpstart_configs,
Expand DownExpand Up@@ -730,10 +730,10 @@ def attach(
ValueError: if the model ID or version cannot be inferred from the training job.

"""

config_name = None
if model_id is None:

model_id, model_version = get_model_id_version_from_training_job(
model_id, model_version, config_name = get_model_info_from_training_job(
training_job_name=training_job_name, sagemaker_session=sagemaker_session
)

Expand All@@ -749,6 +749,7 @@ def attach(
tolerate_deprecated_model=True, # model is already trained, so tolerate if deprecated
tolerate_vulnerable_model=True, # model is already trained, so tolerate if vulnerable
sagemaker_session=sagemaker_session,
config_name=config_name,
)

# eula was already accepted if the model was successfully trained
Expand DownExpand Up@@ -1102,7 +1103,7 @@ def deploy(
tolerate_deprecated_model=self.tolerate_deprecated_model,
tolerate_vulnerable_model=self.tolerate_vulnerable_model,
sagemaker_session=self.sagemaker_session,
# config_name=self.config_name,
config_name=self.config_name,
)

# If a predictor class was passed, do not mutate predictor
Expand Down
5 changes: 4 additions & 1 deletion src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -478,7 +478,10 @@ def _add_tags_to_kwargs(kwargs: JumpStartEstimatorInitKwargs) -> JumpStartEstima

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version
kwargs.tags,
kwargs.model_id,
full_model_version,
config_name=kwargs.config_name,
)
return kwargs

Expand Down
2 changes: 1 addition & 1 deletion src/sagemaker/jumpstart/factory/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -496,7 +496,7 @@ def _add_tags_to_kwargs(kwargs: JumpStartModelDeployKwargs) -> Dict[str, Any]:

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type, kwargs.config_name
)

return kwargs
Expand Down
56 changes: 30 additions & 26 deletions src/sagemaker/jumpstart/session_utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -22,12 +22,12 @@
from sagemaker.utils import aws_partition


def get_model_id_version_from_endpoint(
def get_model_info_from_endpoint(
endpoint_name: str,
inference_component_name: Optional[str] = None,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str, Optional[str]]:
"""Given an endpoint and optionally inference component names, return the model IDand version.
) -> Tuple[str, str, Optional[str], Optional[str]]:
"""Optionally inference component names, return the model ID, version and config name.

Infers the model ID and version based on the resource tags. Returns a tuple of the model ID
and version. A third string element is included in the tuple for any inferred inference
Expand All@@ -46,30 +46,32 @@ def get_model_id_version_from_endpoint(
(
model_id,
model_version,
) = _get_model_id_version_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
config_name,
) = _get_model_info_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
inference_component_name, sagemaker_session
)

else:
(
model_id,
model_version,
config_name,
inference_component_name,
) = _get_model_id_version_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
) = _get_model_info_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
endpoint_name, sagemaker_session
)

else:
model_id, model_version = _get_model_id_version_from_model_based_endpoint(
model_id, model_version, config_name = _get_model_info_from_model_based_endpoint(
endpoint_name, inference_component_name, sagemaker_session
)
return model_id, model_version, inference_component_name
return model_id, model_version, inference_component_name, config_name


def _get_model_id_version_from_inference_component_endpoint_without_inference_component_name(
def _get_model_info_from_inference_component_endpoint_without_inference_component_name(
endpoint_name: str, sagemaker_session: Session
) -> Tuple[str, str, str]:
"""Given an endpoint name, derives the model ID, version, and inferred inference component name.
) -> Tuple[str, str, str, str]:
"""Derives the model ID, version, config name and inferred inference component name.

This function assumes the endpoint corresponds to an inference-component-based endpoint.
An endpoint is inference-component-based if and only if the associated endpoint config
Expand DownExpand Up@@ -98,14 +100,14 @@ def _get_model_id_version_from_inference_component_endpoint_without_inference_co
)
inference_component_name = inference_component_names[0]
return (
*_get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
*_get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name, sagemaker_session
),
inference_component_name,
)


def _get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
def _get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name: str, sagemaker_session: Session
):
"""Returns the model ID and version inferred from a SageMaker inference component.
Expand All@@ -123,7 +125,7 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
f"inference-component/{inference_component_name}"
)

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
inference_component_arn, sagemaker_session
)

Expand All@@ -134,15 +136,15 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
"when retrieving default predictor for this inference component."
)

return model_id, model_version
return model_id, model_version, config_name


def _get_model_id_version_from_model_based_endpoint(
def _get_model_info_from_model_based_endpoint(
endpoint_name: str,
inference_component_name: Optional[str],
sagemaker_session: Session,
) -> Tuple[str, str]:
"""Returns the model IDand version inferred from a model-based endpoint.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID, version and config name inferred from a model-based endpoint.

Raises:
ValueError: If an inference component name is supplied, or if the endpoint does
Expand All@@ -161,7 +163,7 @@ def _get_model_id_version_from_model_based_endpoint(

endpoint_arn = f"arn:{partition}:sagemaker:{region}:{account_id}:endpoint/{endpoint_name}"

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
endpoint_arn, sagemaker_session
)

Expand All@@ -172,14 +174,14 @@ def _get_model_id_version_from_model_based_endpoint(
"predictor for this endpoint."
)

return model_id, model_version
return model_id, model_version, config_name


def get_model_id_version_from_training_job(
def get_model_info_from_training_job(
training_job_name: str,
sagemaker_session: Optional[Session] = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str]:
"""Returns the model ID and version inferred from a training job.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID and version and config name inferred from a training job.

Raises:
ValueError: If the training job does not have tags from which the model ID
Expand All@@ -194,9 +196,11 @@ def get_model_id_version_from_training_job(
f"arn:{partition}:sagemaker:{region}:{account_id}:training-job/{training_job_name}"
)

model_id, inferred_model_version = get_jumpstart_model_id_version_from_resource_arn(
training_job_arn, sagemaker_session
)
(
model_id,
inferred_model_version,
config_name,
) = get_jumpstart_model_id_version_from_resource_arn(training_job_arn, sagemaker_session)

model_version = inferred_model_version or None

Expand All@@ -207,4 +211,4 @@ def get_model_id_version_from_training_job(
"for this training job."
)

return model_id, model_version
return model_id, model_version, config_name
22 changes: 12 additions & 10 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1064,9 +1064,8 @@ def from_json(self, json_obj: Dict[str, Any]) -> None:
Dictionary representation of the config component.
"""
for field in json_obj.keys():
if field not in self.__slots__:
raise ValueError(f"Invalid component field: {field}")
setattr(self, field, json_obj[field])
if field in self.__slots__:
setattr(self, field, json_obj[field])


class JumpStartMetadataConfig(JumpStartDataHolderType):
Expand DownExpand Up@@ -1164,20 +1163,17 @@ def get_top_config_from_ranking(
) -> Optional[JumpStartMetadataConfig]:
"""Gets the best the config based on config ranking.

Fallback to use the ordering in config names if
ranking is not available.
Args:
ranking_name (str):
The ranking name that config priority is based on.
instance_type (Optional[str]):
The instance type which the config selection is based on.

Raises:
ValueError: If the config exists but missing config ranking.
NotImplementedError: If the scope is unrecognized.
"""
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
raise ValueError(f"Config exists but missing config ranking {ranking_name}.")

if self.scope == JumpStartScriptScope.INFERENCE:
instance_type_attribute = "supported_inference_instance_types"
Expand All@@ -1186,8 +1182,14 @@ def get_top_config_from_ranking(
else:
raise NotImplementedError(f"Unknown script scope {self.scope}")

rankings = self.config_rankings.get(ranking_name)
for config_name in rankings.rankings:
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
ranked_config_names = sorted(list(self.configs.keys()))
else:
rankings = self.config_rankings.get(ranking_name)
ranked_config_names = rankings.rankings
for config_name in ranked_config_names:
resolved_config = self.configs[config_name].resolved_config
if instance_type and instance_type not in getattr(
resolved_config, instance_type_attribute
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
Merged
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
1 change: 1 addition & 0 deletions src/sagemaker/jumpstart/enums.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -92,6 +92,7 @@ class JumpStartTag(str, Enum):
MODEL_ID = "sagemaker-sdk:jumpstart-model-id"
MODEL_VERSION = "sagemaker-sdk:jumpstart-model-version"
MODEL_TYPE = "sagemaker-sdk:jumpstart-model-type"
MODEL_CONFIG_NAME = "sagemaker-sdk:jumpstart-model-config-name"


class SerializerType(str, Enum):
Expand Down
9 changes: 5 additions & 4 deletions src/sagemaker/jumpstart/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -33,7 +33,7 @@

from sagemaker.jumpstart.factory.estimator import get_deploy_kwargs, get_fit_kwargs, get_init_kwargs
from sagemaker.jumpstart.factory.model import get_default_predictor
from sagemaker.jumpstart.session_utils import get_model_id_version_from_training_job
from sagemaker.jumpstart.session_utils import get_model_info_from_training_job
from sagemaker.jumpstart.types import JumpStartMetadataConfig
from sagemaker.jumpstart.utils import (
get_jumpstart_configs,
Expand DownExpand Up@@ -730,10 +730,10 @@ def attach(
ValueError: if the model ID or version cannot be inferred from the training job.

"""

config_name = None
if model_id is None:

model_id, model_version = get_model_id_version_from_training_job(
model_id, model_version, config_name = get_model_info_from_training_job(
training_job_name=training_job_name, sagemaker_session=sagemaker_session
)

Expand All@@ -749,6 +749,7 @@ def attach(
tolerate_deprecated_model=True, # model is already trained, so tolerate if deprecated
tolerate_vulnerable_model=True, # model is already trained, so tolerate if vulnerable
sagemaker_session=sagemaker_session,
config_name=config_name,
)

# eula was already accepted if the model was successfully trained
Expand DownExpand Up@@ -1102,7 +1103,7 @@ def deploy(
tolerate_deprecated_model=self.tolerate_deprecated_model,
tolerate_vulnerable_model=self.tolerate_vulnerable_model,
sagemaker_session=self.sagemaker_session,
# config_name=self.config_name,
config_name=self.config_name,
)

# If a predictor class was passed, do not mutate predictor
Expand Down
5 changes: 4 additions & 1 deletion src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -478,7 +478,10 @@ def _add_tags_to_kwargs(kwargs: JumpStartEstimatorInitKwargs) -> JumpStartEstima

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version
kwargs.tags,
kwargs.model_id,
full_model_version,
config_name=kwargs.config_name,
)
return kwargs

Expand Down
2 changes: 1 addition & 1 deletion src/sagemaker/jumpstart/factory/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -496,7 +496,7 @@ def _add_tags_to_kwargs(kwargs: JumpStartModelDeployKwargs) -> Dict[str, Any]:

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type, kwargs.config_name
)

return kwargs
Expand Down
56 changes: 30 additions & 26 deletions src/sagemaker/jumpstart/session_utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -22,12 +22,12 @@
from sagemaker.utils import aws_partition


def get_model_id_version_from_endpoint(
def get_model_info_from_endpoint(
endpoint_name: str,
inference_component_name: Optional[str] = None,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str, Optional[str]]:
"""Given an endpoint and optionally inference component names, return the model IDand version.
) -> Tuple[str, str, Optional[str], Optional[str]]:
"""Optionally inference component names, return the model ID, version and config name.

Infers the model ID and version based on the resource tags. Returns a tuple of the model ID
and version. A third string element is included in the tuple for any inferred inference
Expand All@@ -46,30 +46,32 @@ def get_model_id_version_from_endpoint(
(
model_id,
model_version,
) = _get_model_id_version_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
config_name,
) = _get_model_info_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
inference_component_name, sagemaker_session
)

else:
(
model_id,
model_version,
config_name,
inference_component_name,
) = _get_model_id_version_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
) = _get_model_info_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
endpoint_name, sagemaker_session
)

else:
model_id, model_version = _get_model_id_version_from_model_based_endpoint(
model_id, model_version, config_name = _get_model_info_from_model_based_endpoint(
endpoint_name, inference_component_name, sagemaker_session
)
return model_id, model_version, inference_component_name
return model_id, model_version, inference_component_name, config_name


def _get_model_id_version_from_inference_component_endpoint_without_inference_component_name(
def _get_model_info_from_inference_component_endpoint_without_inference_component_name(
endpoint_name: str, sagemaker_session: Session
) -> Tuple[str, str, str]:
"""Given an endpoint name, derives the model ID, version, and inferred inference component name.
) -> Tuple[str, str, str, str]:
"""Derives the model ID, version, config name and inferred inference component name.

This function assumes the endpoint corresponds to an inference-component-based endpoint.
An endpoint is inference-component-based if and only if the associated endpoint config
Expand DownExpand Up@@ -98,14 +100,14 @@ def _get_model_id_version_from_inference_component_endpoint_without_inference_co
)
inference_component_name = inference_component_names[0]
return (
*_get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
*_get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name, sagemaker_session
),
inference_component_name,
)


def _get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
def _get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name: str, sagemaker_session: Session
):
"""Returns the model ID and version inferred from a SageMaker inference component.
Expand All@@ -123,7 +125,7 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
f"inference-component/{inference_component_name}"
)

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
inference_component_arn, sagemaker_session
)

Expand All@@ -134,15 +136,15 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
"when retrieving default predictor for this inference component."
)

return model_id, model_version
return model_id, model_version, config_name


def _get_model_id_version_from_model_based_endpoint(
def _get_model_info_from_model_based_endpoint(
endpoint_name: str,
inference_component_name: Optional[str],
sagemaker_session: Session,
) -> Tuple[str, str]:
"""Returns the model IDand version inferred from a model-based endpoint.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID, version and config name inferred from a model-based endpoint.

Raises:
ValueError: If an inference component name is supplied, or if the endpoint does
Expand All@@ -161,7 +163,7 @@ def _get_model_id_version_from_model_based_endpoint(

endpoint_arn = f"arn:{partition}:sagemaker:{region}:{account_id}:endpoint/{endpoint_name}"

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
endpoint_arn, sagemaker_session
)

Expand All@@ -172,14 +174,14 @@ def _get_model_id_version_from_model_based_endpoint(
"predictor for this endpoint."
)

return model_id, model_version
return model_id, model_version, config_name


def get_model_id_version_from_training_job(
def get_model_info_from_training_job(
training_job_name: str,
sagemaker_session: Optional[Session] = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str]:
"""Returns the model ID and version inferred from a training job.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID and version and config name inferred from a training job.

Raises:
ValueError: If the training job does not have tags from which the model ID
Expand All@@ -194,9 +196,11 @@ def get_model_id_version_from_training_job(
f"arn:{partition}:sagemaker:{region}:{account_id}:training-job/{training_job_name}"
)

model_id, inferred_model_version = get_jumpstart_model_id_version_from_resource_arn(
training_job_arn, sagemaker_session
)
(
model_id,
inferred_model_version,
config_name,
) = get_jumpstart_model_id_version_from_resource_arn(training_job_arn, sagemaker_session)

model_version = inferred_model_version or None

Expand All@@ -207,4 +211,4 @@ def get_model_id_version_from_training_job(
"for this training job."
)

return model_id, model_version
return model_id, model_version, config_name
22 changes: 12 additions & 10 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1064,9 +1064,8 @@ def from_json(self, json_obj: Dict[str, Any]) -> None:
Dictionary representation of the config component.
"""
for field in json_obj.keys():
if field not in self.__slots__:
raise ValueError(f"Invalid component field: {field}")
setattr(self, field, json_obj[field])
if field in self.__slots__:
setattr(self, field, json_obj[field])


class JumpStartMetadataConfig(JumpStartDataHolderType):
Expand DownExpand Up@@ -1164,20 +1163,17 @@ def get_top_config_from_ranking(
) -> Optional[JumpStartMetadataConfig]:
"""Gets the best the config based on config ranking.

Fallback to use the ordering in config names if
ranking is not available.
Args:
ranking_name (str):
The ranking name that config priority is based on.
instance_type (Optional[str]):
The instance type which the config selection is based on.

Raises:
ValueError: If the config exists but missing config ranking.
NotImplementedError: If the scope is unrecognized.
"""
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
raise ValueError(f"Config exists but missing config ranking {ranking_name}.")

if self.scope == JumpStartScriptScope.INFERENCE:
instance_type_attribute = "supported_inference_instance_types"
Expand All@@ -1186,8 +1182,14 @@ def get_top_config_from_ranking(
else:
raise NotImplementedError(f"Unknown script scope {self.scope}")

rankings = self.config_rankings.get(ranking_name)
for config_name in rankings.rankings:
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
ranked_config_names = sorted(list(self.configs.keys()))
else:
rankings = self.config_rankings.get(ranking_name)
ranked_config_names = rankings.rankings
for config_name in ranked_config_names:
resolved_config = self.configs[config_name].resolved_config
if instance_type and instance_type not in getattr(
resolved_config, instance_type_attribute
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
Merged
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
1 change: 1 addition & 0 deletions src/sagemaker/jumpstart/enums.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -92,6 +92,7 @@ class JumpStartTag(str, Enum):
MODEL_ID = "sagemaker-sdk:jumpstart-model-id"
MODEL_VERSION = "sagemaker-sdk:jumpstart-model-version"
MODEL_TYPE = "sagemaker-sdk:jumpstart-model-type"
MODEL_CONFIG_NAME = "sagemaker-sdk:jumpstart-model-config-name"


class SerializerType(str, Enum):
Expand Down
9 changes: 5 additions & 4 deletions src/sagemaker/jumpstart/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -33,7 +33,7 @@

from sagemaker.jumpstart.factory.estimator import get_deploy_kwargs, get_fit_kwargs, get_init_kwargs
from sagemaker.jumpstart.factory.model import get_default_predictor
from sagemaker.jumpstart.session_utils import get_model_id_version_from_training_job
from sagemaker.jumpstart.session_utils import get_model_info_from_training_job
from sagemaker.jumpstart.types import JumpStartMetadataConfig
from sagemaker.jumpstart.utils import (
get_jumpstart_configs,
Expand DownExpand Up@@ -730,10 +730,10 @@ def attach(
ValueError: if the model ID or version cannot be inferred from the training job.

"""

config_name = None
if model_id is None:

model_id, model_version = get_model_id_version_from_training_job(
model_id, model_version, config_name = get_model_info_from_training_job(
training_job_name=training_job_name, sagemaker_session=sagemaker_session
)

Expand All@@ -749,6 +749,7 @@ def attach(
tolerate_deprecated_model=True, # model is already trained, so tolerate if deprecated
tolerate_vulnerable_model=True, # model is already trained, so tolerate if vulnerable
sagemaker_session=sagemaker_session,
config_name=config_name,
)

# eula was already accepted if the model was successfully trained
Expand DownExpand Up@@ -1102,7 +1103,7 @@ def deploy(
tolerate_deprecated_model=self.tolerate_deprecated_model,
tolerate_vulnerable_model=self.tolerate_vulnerable_model,
sagemaker_session=self.sagemaker_session,
# config_name=self.config_name,
config_name=self.config_name,
)

# If a predictor class was passed, do not mutate predictor
Expand Down
5 changes: 4 additions & 1 deletion src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -478,7 +478,10 @@ def _add_tags_to_kwargs(kwargs: JumpStartEstimatorInitKwargs) -> JumpStartEstima

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version
kwargs.tags,
kwargs.model_id,
full_model_version,
config_name=kwargs.config_name,
)
return kwargs

Expand Down
2 changes: 1 addition & 1 deletion src/sagemaker/jumpstart/factory/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -496,7 +496,7 @@ def _add_tags_to_kwargs(kwargs: JumpStartModelDeployKwargs) -> Dict[str, Any]:

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type, kwargs.config_name
)

return kwargs
Expand Down
56 changes: 30 additions & 26 deletions src/sagemaker/jumpstart/session_utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -22,12 +22,12 @@
from sagemaker.utils import aws_partition


def get_model_id_version_from_endpoint(
def get_model_info_from_endpoint(
endpoint_name: str,
inference_component_name: Optional[str] = None,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str, Optional[str]]:
"""Given an endpoint and optionally inference component names, return the model IDand version.
) -> Tuple[str, str, Optional[str], Optional[str]]:
"""Optionally inference component names, return the model ID, version and config name.

Infers the model ID and version based on the resource tags. Returns a tuple of the model ID
and version. A third string element is included in the tuple for any inferred inference
Expand All@@ -46,30 +46,32 @@ def get_model_id_version_from_endpoint(
(
model_id,
model_version,
) = _get_model_id_version_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
config_name,
) = _get_model_info_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
inference_component_name, sagemaker_session
)

else:
(
model_id,
model_version,
config_name,
inference_component_name,
) = _get_model_id_version_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
) = _get_model_info_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
endpoint_name, sagemaker_session
)

else:
model_id, model_version = _get_model_id_version_from_model_based_endpoint(
model_id, model_version, config_name = _get_model_info_from_model_based_endpoint(
endpoint_name, inference_component_name, sagemaker_session
)
return model_id, model_version, inference_component_name
return model_id, model_version, inference_component_name, config_name


def _get_model_id_version_from_inference_component_endpoint_without_inference_component_name(
def _get_model_info_from_inference_component_endpoint_without_inference_component_name(
endpoint_name: str, sagemaker_session: Session
) -> Tuple[str, str, str]:
"""Given an endpoint name, derives the model ID, version, and inferred inference component name.
) -> Tuple[str, str, str, str]:
"""Derives the model ID, version, config name and inferred inference component name.

This function assumes the endpoint corresponds to an inference-component-based endpoint.
An endpoint is inference-component-based if and only if the associated endpoint config
Expand DownExpand Up@@ -98,14 +100,14 @@ def _get_model_id_version_from_inference_component_endpoint_without_inference_co
)
inference_component_name = inference_component_names[0]
return (
*_get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
*_get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name, sagemaker_session
),
inference_component_name,
)


def _get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
def _get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name: str, sagemaker_session: Session
):
"""Returns the model ID and version inferred from a SageMaker inference component.
Expand All@@ -123,7 +125,7 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
f"inference-component/{inference_component_name}"
)

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
inference_component_arn, sagemaker_session
)

Expand All@@ -134,15 +136,15 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
"when retrieving default predictor for this inference component."
)

return model_id, model_version
return model_id, model_version, config_name


def _get_model_id_version_from_model_based_endpoint(
def _get_model_info_from_model_based_endpoint(
endpoint_name: str,
inference_component_name: Optional[str],
sagemaker_session: Session,
) -> Tuple[str, str]:
"""Returns the model IDand version inferred from a model-based endpoint.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID, version and config name inferred from a model-based endpoint.

Raises:
ValueError: If an inference component name is supplied, or if the endpoint does
Expand All@@ -161,7 +163,7 @@ def _get_model_id_version_from_model_based_endpoint(

endpoint_arn = f"arn:{partition}:sagemaker:{region}:{account_id}:endpoint/{endpoint_name}"

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
endpoint_arn, sagemaker_session
)

Expand All@@ -172,14 +174,14 @@ def _get_model_id_version_from_model_based_endpoint(
"predictor for this endpoint."
)

return model_id, model_version
return model_id, model_version, config_name


def get_model_id_version_from_training_job(
def get_model_info_from_training_job(
training_job_name: str,
sagemaker_session: Optional[Session] = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str]:
"""Returns the model ID and version inferred from a training job.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID and version and config name inferred from a training job.

Raises:
ValueError: If the training job does not have tags from which the model ID
Expand All@@ -194,9 +196,11 @@ def get_model_id_version_from_training_job(
f"arn:{partition}:sagemaker:{region}:{account_id}:training-job/{training_job_name}"
)

model_id, inferred_model_version = get_jumpstart_model_id_version_from_resource_arn(
training_job_arn, sagemaker_session
)
(
model_id,
inferred_model_version,
config_name,
) = get_jumpstart_model_id_version_from_resource_arn(training_job_arn, sagemaker_session)

model_version = inferred_model_version or None

Expand All@@ -207,4 +211,4 @@ def get_model_id_version_from_training_job(
"for this training job."
)

return model_id, model_version
return model_id, model_version, config_name
22 changes: 12 additions & 10 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1064,9 +1064,8 @@ def from_json(self, json_obj: Dict[str, Any]) -> None:
Dictionary representation of the config component.
"""
for field in json_obj.keys():
if field not in self.__slots__:
raise ValueError(f"Invalid component field: {field}")
setattr(self, field, json_obj[field])
if field in self.__slots__:
setattr(self, field, json_obj[field])


class JumpStartMetadataConfig(JumpStartDataHolderType):
Expand DownExpand Up@@ -1164,20 +1163,17 @@ def get_top_config_from_ranking(
) -> Optional[JumpStartMetadataConfig]:
"""Gets the best the config based on config ranking.

Fallback to use the ordering in config names if
ranking is not available.
Args:
ranking_name (str):
The ranking name that config priority is based on.
instance_type (Optional[str]):
The instance type which the config selection is based on.

Raises:
ValueError: If the config exists but missing config ranking.
NotImplementedError: If the scope is unrecognized.
"""
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
raise ValueError(f"Config exists but missing config ranking {ranking_name}.")

if self.scope == JumpStartScriptScope.INFERENCE:
instance_type_attribute = "supported_inference_instance_types"
Expand All@@ -1186,8 +1182,14 @@ def get_top_config_from_ranking(
else:
raise NotImplementedError(f"Unknown script scope {self.scope}")

rankings = self.config_rankings.get(ranking_name)
for config_name in rankings.rankings:
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
ranked_config_names = sorted(list(self.configs.keys()))
else:
rankings = self.config_rankings.get(ranking_name)
ranked_config_names = rankings.rankings
for config_name in ranked_config_names:
resolved_config = self.configs[config_name].resolved_config
if instance_type and instance_type not in getattr(
resolved_config, instance_type_attribute
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
Merged
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
1 change: 1 addition & 0 deletions src/sagemaker/jumpstart/enums.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -92,6 +92,7 @@ class JumpStartTag(str, Enum):
MODEL_ID = "sagemaker-sdk:jumpstart-model-id"
MODEL_VERSION = "sagemaker-sdk:jumpstart-model-version"
MODEL_TYPE = "sagemaker-sdk:jumpstart-model-type"
MODEL_CONFIG_NAME = "sagemaker-sdk:jumpstart-model-config-name"


class SerializerType(str, Enum):
Expand Down
9 changes: 5 additions & 4 deletions src/sagemaker/jumpstart/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -33,7 +33,7 @@

from sagemaker.jumpstart.factory.estimator import get_deploy_kwargs, get_fit_kwargs, get_init_kwargs
from sagemaker.jumpstart.factory.model import get_default_predictor
from sagemaker.jumpstart.session_utils import get_model_id_version_from_training_job
from sagemaker.jumpstart.session_utils import get_model_info_from_training_job
from sagemaker.jumpstart.types import JumpStartMetadataConfig
from sagemaker.jumpstart.utils import (
get_jumpstart_configs,
Expand DownExpand Up@@ -730,10 +730,10 @@ def attach(
ValueError: if the model ID or version cannot be inferred from the training job.

"""

config_name = None
if model_id is None:

model_id, model_version = get_model_id_version_from_training_job(
model_id, model_version, config_name = get_model_info_from_training_job(
training_job_name=training_job_name, sagemaker_session=sagemaker_session
)

Expand All@@ -749,6 +749,7 @@ def attach(
tolerate_deprecated_model=True, # model is already trained, so tolerate if deprecated
tolerate_vulnerable_model=True, # model is already trained, so tolerate if vulnerable
sagemaker_session=sagemaker_session,
config_name=config_name,
)

# eula was already accepted if the model was successfully trained
Expand DownExpand Up@@ -1102,7 +1103,7 @@ def deploy(
tolerate_deprecated_model=self.tolerate_deprecated_model,
tolerate_vulnerable_model=self.tolerate_vulnerable_model,
sagemaker_session=self.sagemaker_session,
# config_name=self.config_name,
config_name=self.config_name,
)

# If a predictor class was passed, do not mutate predictor
Expand Down
5 changes: 4 additions & 1 deletion src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -478,7 +478,10 @@ def _add_tags_to_kwargs(kwargs: JumpStartEstimatorInitKwargs) -> JumpStartEstima

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version
kwargs.tags,
kwargs.model_id,
full_model_version,
config_name=kwargs.config_name,
)
return kwargs

Expand Down
2 changes: 1 addition & 1 deletion src/sagemaker/jumpstart/factory/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -496,7 +496,7 @@ def _add_tags_to_kwargs(kwargs: JumpStartModelDeployKwargs) -> Dict[str, Any]:

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type, kwargs.config_name
)

return kwargs
Expand Down
56 changes: 30 additions & 26 deletions src/sagemaker/jumpstart/session_utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -22,12 +22,12 @@
from sagemaker.utils import aws_partition


def get_model_id_version_from_endpoint(
def get_model_info_from_endpoint(
endpoint_name: str,
inference_component_name: Optional[str] = None,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str, Optional[str]]:
"""Given an endpoint and optionally inference component names, return the model IDand version.
) -> Tuple[str, str, Optional[str], Optional[str]]:
"""Optionally inference component names, return the model ID, version and config name.

Infers the model ID and version based on the resource tags. Returns a tuple of the model ID
and version. A third string element is included in the tuple for any inferred inference
Expand All@@ -46,30 +46,32 @@ def get_model_id_version_from_endpoint(
(
model_id,
model_version,
) = _get_model_id_version_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
config_name,
) = _get_model_info_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
inference_component_name, sagemaker_session
)

else:
(
model_id,
model_version,
config_name,
inference_component_name,
) = _get_model_id_version_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
) = _get_model_info_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
endpoint_name, sagemaker_session
)

else:
model_id, model_version = _get_model_id_version_from_model_based_endpoint(
model_id, model_version, config_name = _get_model_info_from_model_based_endpoint(
endpoint_name, inference_component_name, sagemaker_session
)
return model_id, model_version, inference_component_name
return model_id, model_version, inference_component_name, config_name


def _get_model_id_version_from_inference_component_endpoint_without_inference_component_name(
def _get_model_info_from_inference_component_endpoint_without_inference_component_name(
endpoint_name: str, sagemaker_session: Session
) -> Tuple[str, str, str]:
"""Given an endpoint name, derives the model ID, version, and inferred inference component name.
) -> Tuple[str, str, str, str]:
"""Derives the model ID, version, config name and inferred inference component name.

This function assumes the endpoint corresponds to an inference-component-based endpoint.
An endpoint is inference-component-based if and only if the associated endpoint config
Expand DownExpand Up@@ -98,14 +100,14 @@ def _get_model_id_version_from_inference_component_endpoint_without_inference_co
)
inference_component_name = inference_component_names[0]
return (
*_get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
*_get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name, sagemaker_session
),
inference_component_name,
)


def _get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
def _get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name: str, sagemaker_session: Session
):
"""Returns the model ID and version inferred from a SageMaker inference component.
Expand All@@ -123,7 +125,7 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
f"inference-component/{inference_component_name}"
)

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
inference_component_arn, sagemaker_session
)

Expand All@@ -134,15 +136,15 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
"when retrieving default predictor for this inference component."
)

return model_id, model_version
return model_id, model_version, config_name


def _get_model_id_version_from_model_based_endpoint(
def _get_model_info_from_model_based_endpoint(
endpoint_name: str,
inference_component_name: Optional[str],
sagemaker_session: Session,
) -> Tuple[str, str]:
"""Returns the model IDand version inferred from a model-based endpoint.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID, version and config name inferred from a model-based endpoint.

Raises:
ValueError: If an inference component name is supplied, or if the endpoint does
Expand All@@ -161,7 +163,7 @@ def _get_model_id_version_from_model_based_endpoint(

endpoint_arn = f"arn:{partition}:sagemaker:{region}:{account_id}:endpoint/{endpoint_name}"

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
endpoint_arn, sagemaker_session
)

Expand All@@ -172,14 +174,14 @@ def _get_model_id_version_from_model_based_endpoint(
"predictor for this endpoint."
)

return model_id, model_version
return model_id, model_version, config_name


def get_model_id_version_from_training_job(
def get_model_info_from_training_job(
training_job_name: str,
sagemaker_session: Optional[Session] = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str]:
"""Returns the model ID and version inferred from a training job.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID and version and config name inferred from a training job.

Raises:
ValueError: If the training job does not have tags from which the model ID
Expand All@@ -194,9 +196,11 @@ def get_model_id_version_from_training_job(
f"arn:{partition}:sagemaker:{region}:{account_id}:training-job/{training_job_name}"
)

model_id, inferred_model_version = get_jumpstart_model_id_version_from_resource_arn(
training_job_arn, sagemaker_session
)
(
model_id,
inferred_model_version,
config_name,
) = get_jumpstart_model_id_version_from_resource_arn(training_job_arn, sagemaker_session)

model_version = inferred_model_version or None

Expand All@@ -207,4 +211,4 @@ def get_model_id_version_from_training_job(
"for this training job."
)

return model_id, model_version
return model_id, model_version, config_name
22 changes: 12 additions & 10 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1064,9 +1064,8 @@ def from_json(self, json_obj: Dict[str, Any]) -> None:
Dictionary representation of the config component.
"""
for field in json_obj.keys():
if field not in self.__slots__:
raise ValueError(f"Invalid component field: {field}")
setattr(self, field, json_obj[field])
if field in self.__slots__:
setattr(self, field, json_obj[field])


class JumpStartMetadataConfig(JumpStartDataHolderType):
Expand DownExpand Up@@ -1164,20 +1163,17 @@ def get_top_config_from_ranking(
) -> Optional[JumpStartMetadataConfig]:
"""Gets the best the config based on config ranking.

Fallback to use the ordering in config names if
ranking is not available.
Args:
ranking_name (str):
The ranking name that config priority is based on.
instance_type (Optional[str]):
The instance type which the config selection is based on.

Raises:
ValueError: If the config exists but missing config ranking.
NotImplementedError: If the scope is unrecognized.
"""
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
raise ValueError(f"Config exists but missing config ranking {ranking_name}.")

if self.scope == JumpStartScriptScope.INFERENCE:
instance_type_attribute = "supported_inference_instance_types"
Expand All@@ -1186,8 +1182,14 @@ def get_top_config_from_ranking(
else:
raise NotImplementedError(f"Unknown script scope {self.scope}")

rankings = self.config_rankings.get(ranking_name)
for config_name in rankings.rankings:
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
ranked_config_names = sorted(list(self.configs.keys()))
else:
rankings = self.config_rankings.get(ranking_name)
ranked_config_names = rankings.rankings
for config_name in ranked_config_names:
resolved_config = self.configs[config_name].resolved_config
if instance_type and instance_type not in getattr(
resolved_config, instance_type_attribute
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
Merged
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
1 change: 1 addition & 0 deletions src/sagemaker/jumpstart/enums.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -92,6 +92,7 @@ class JumpStartTag(str, Enum):
MODEL_ID = "sagemaker-sdk:jumpstart-model-id"
MODEL_VERSION = "sagemaker-sdk:jumpstart-model-version"
MODEL_TYPE = "sagemaker-sdk:jumpstart-model-type"
MODEL_CONFIG_NAME = "sagemaker-sdk:jumpstart-model-config-name"


class SerializerType(str, Enum):
Expand Down
9 changes: 5 additions & 4 deletions src/sagemaker/jumpstart/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -33,7 +33,7 @@

from sagemaker.jumpstart.factory.estimator import get_deploy_kwargs, get_fit_kwargs, get_init_kwargs
from sagemaker.jumpstart.factory.model import get_default_predictor
from sagemaker.jumpstart.session_utils import get_model_id_version_from_training_job
from sagemaker.jumpstart.session_utils import get_model_info_from_training_job
from sagemaker.jumpstart.types import JumpStartMetadataConfig
from sagemaker.jumpstart.utils import (
get_jumpstart_configs,
Expand DownExpand Up@@ -730,10 +730,10 @@ def attach(
ValueError: if the model ID or version cannot be inferred from the training job.

"""

config_name = None
if model_id is None:

model_id, model_version = get_model_id_version_from_training_job(
model_id, model_version, config_name = get_model_info_from_training_job(
training_job_name=training_job_name, sagemaker_session=sagemaker_session
)

Expand All@@ -749,6 +749,7 @@ def attach(
tolerate_deprecated_model=True, # model is already trained, so tolerate if deprecated
tolerate_vulnerable_model=True, # model is already trained, so tolerate if vulnerable
sagemaker_session=sagemaker_session,
config_name=config_name,
)

# eula was already accepted if the model was successfully trained
Expand DownExpand Up@@ -1102,7 +1103,7 @@ def deploy(
tolerate_deprecated_model=self.tolerate_deprecated_model,
tolerate_vulnerable_model=self.tolerate_vulnerable_model,
sagemaker_session=self.sagemaker_session,
# config_name=self.config_name,
config_name=self.config_name,
)

# If a predictor class was passed, do not mutate predictor
Expand Down
5 changes: 4 additions & 1 deletion src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -478,7 +478,10 @@ def _add_tags_to_kwargs(kwargs: JumpStartEstimatorInitKwargs) -> JumpStartEstima

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version
kwargs.tags,
kwargs.model_id,
full_model_version,
config_name=kwargs.config_name,
)
return kwargs

Expand Down
2 changes: 1 addition & 1 deletion src/sagemaker/jumpstart/factory/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -496,7 +496,7 @@ def _add_tags_to_kwargs(kwargs: JumpStartModelDeployKwargs) -> Dict[str, Any]:

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type, kwargs.config_name
)

return kwargs
Expand Down
56 changes: 30 additions & 26 deletions src/sagemaker/jumpstart/session_utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -22,12 +22,12 @@
from sagemaker.utils import aws_partition


def get_model_id_version_from_endpoint(
def get_model_info_from_endpoint(
endpoint_name: str,
inference_component_name: Optional[str] = None,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str, Optional[str]]:
"""Given an endpoint and optionally inference component names, return the model IDand version.
) -> Tuple[str, str, Optional[str], Optional[str]]:
"""Optionally inference component names, return the model ID, version and config name.

Infers the model ID and version based on the resource tags. Returns a tuple of the model ID
and version. A third string element is included in the tuple for any inferred inference
Expand All@@ -46,30 +46,32 @@ def get_model_id_version_from_endpoint(
(
model_id,
model_version,
) = _get_model_id_version_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
config_name,
) = _get_model_info_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
inference_component_name, sagemaker_session
)

else:
(
model_id,
model_version,
config_name,
inference_component_name,
) = _get_model_id_version_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
) = _get_model_info_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
endpoint_name, sagemaker_session
)

else:
model_id, model_version = _get_model_id_version_from_model_based_endpoint(
model_id, model_version, config_name = _get_model_info_from_model_based_endpoint(
endpoint_name, inference_component_name, sagemaker_session
)
return model_id, model_version, inference_component_name
return model_id, model_version, inference_component_name, config_name


def _get_model_id_version_from_inference_component_endpoint_without_inference_component_name(
def _get_model_info_from_inference_component_endpoint_without_inference_component_name(
endpoint_name: str, sagemaker_session: Session
) -> Tuple[str, str, str]:
"""Given an endpoint name, derives the model ID, version, and inferred inference component name.
) -> Tuple[str, str, str, str]:
"""Derives the model ID, version, config name and inferred inference component name.

This function assumes the endpoint corresponds to an inference-component-based endpoint.
An endpoint is inference-component-based if and only if the associated endpoint config
Expand DownExpand Up@@ -98,14 +100,14 @@ def _get_model_id_version_from_inference_component_endpoint_without_inference_co
)
inference_component_name = inference_component_names[0]
return (
*_get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
*_get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name, sagemaker_session
),
inference_component_name,
)


def _get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
def _get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name: str, sagemaker_session: Session
):
"""Returns the model ID and version inferred from a SageMaker inference component.
Expand All@@ -123,7 +125,7 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
f"inference-component/{inference_component_name}"
)

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
inference_component_arn, sagemaker_session
)

Expand All@@ -134,15 +136,15 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
"when retrieving default predictor for this inference component."
)

return model_id, model_version
return model_id, model_version, config_name


def _get_model_id_version_from_model_based_endpoint(
def _get_model_info_from_model_based_endpoint(
endpoint_name: str,
inference_component_name: Optional[str],
sagemaker_session: Session,
) -> Tuple[str, str]:
"""Returns the model IDand version inferred from a model-based endpoint.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID, version and config name inferred from a model-based endpoint.

Raises:
ValueError: If an inference component name is supplied, or if the endpoint does
Expand All@@ -161,7 +163,7 @@ def _get_model_id_version_from_model_based_endpoint(

endpoint_arn = f"arn:{partition}:sagemaker:{region}:{account_id}:endpoint/{endpoint_name}"

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
endpoint_arn, sagemaker_session
)

Expand All@@ -172,14 +174,14 @@ def _get_model_id_version_from_model_based_endpoint(
"predictor for this endpoint."
)

return model_id, model_version
return model_id, model_version, config_name


def get_model_id_version_from_training_job(
def get_model_info_from_training_job(
training_job_name: str,
sagemaker_session: Optional[Session] = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str]:
"""Returns the model ID and version inferred from a training job.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID and version and config name inferred from a training job.

Raises:
ValueError: If the training job does not have tags from which the model ID
Expand All@@ -194,9 +196,11 @@ def get_model_id_version_from_training_job(
f"arn:{partition}:sagemaker:{region}:{account_id}:training-job/{training_job_name}"
)

model_id, inferred_model_version = get_jumpstart_model_id_version_from_resource_arn(
training_job_arn, sagemaker_session
)
(
model_id,
inferred_model_version,
config_name,
) = get_jumpstart_model_id_version_from_resource_arn(training_job_arn, sagemaker_session)

model_version = inferred_model_version or None

Expand All@@ -207,4 +211,4 @@ def get_model_id_version_from_training_job(
"for this training job."
)

return model_id, model_version
return model_id, model_version, config_name
22 changes: 12 additions & 10 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1064,9 +1064,8 @@ def from_json(self, json_obj: Dict[str, Any]) -> None:
Dictionary representation of the config component.
"""
for field in json_obj.keys():
if field not in self.__slots__:
raise ValueError(f"Invalid component field: {field}")
setattr(self, field, json_obj[field])
if field in self.__slots__:
setattr(self, field, json_obj[field])


class JumpStartMetadataConfig(JumpStartDataHolderType):
Expand DownExpand Up@@ -1164,20 +1163,17 @@ def get_top_config_from_ranking(
) -> Optional[JumpStartMetadataConfig]:
"""Gets the best the config based on config ranking.

Fallback to use the ordering in config names if
ranking is not available.
Args:
ranking_name (str):
The ranking name that config priority is based on.
instance_type (Optional[str]):
The instance type which the config selection is based on.

Raises:
ValueError: If the config exists but missing config ranking.
NotImplementedError: If the scope is unrecognized.
"""
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
raise ValueError(f"Config exists but missing config ranking {ranking_name}.")

if self.scope == JumpStartScriptScope.INFERENCE:
instance_type_attribute = "supported_inference_instance_types"
Expand All@@ -1186,8 +1182,14 @@ def get_top_config_from_ranking(
else:
raise NotImplementedError(f"Unknown script scope {self.scope}")

rankings = self.config_rankings.get(ranking_name)
for config_name in rankings.rankings:
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
ranked_config_names = sorted(list(self.configs.keys()))
else:
rankings = self.config_rankings.get(ranking_name)
ranked_config_names = rankings.rankings
for config_name in ranked_config_names:
resolved_config = self.configs[config_name].resolved_config
if instance_type and instance_type not in getattr(
resolved_config, instance_type_attribute
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
Merged
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
1 change: 1 addition & 0 deletions src/sagemaker/jumpstart/enums.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -92,6 +92,7 @@ class JumpStartTag(str, Enum):
MODEL_ID = "sagemaker-sdk:jumpstart-model-id"
MODEL_VERSION = "sagemaker-sdk:jumpstart-model-version"
MODEL_TYPE = "sagemaker-sdk:jumpstart-model-type"
MODEL_CONFIG_NAME = "sagemaker-sdk:jumpstart-model-config-name"


class SerializerType(str, Enum):
Expand Down
9 changes: 5 additions & 4 deletions src/sagemaker/jumpstart/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -33,7 +33,7 @@

from sagemaker.jumpstart.factory.estimator import get_deploy_kwargs, get_fit_kwargs, get_init_kwargs
from sagemaker.jumpstart.factory.model import get_default_predictor
from sagemaker.jumpstart.session_utils import get_model_id_version_from_training_job
from sagemaker.jumpstart.session_utils import get_model_info_from_training_job
from sagemaker.jumpstart.types import JumpStartMetadataConfig
from sagemaker.jumpstart.utils import (
get_jumpstart_configs,
Expand DownExpand Up@@ -730,10 +730,10 @@ def attach(
ValueError: if the model ID or version cannot be inferred from the training job.

"""

config_name = None
if model_id is None:

model_id, model_version = get_model_id_version_from_training_job(
model_id, model_version, config_name = get_model_info_from_training_job(
training_job_name=training_job_name, sagemaker_session=sagemaker_session
)

Expand All@@ -749,6 +749,7 @@ def attach(
tolerate_deprecated_model=True, # model is already trained, so tolerate if deprecated
tolerate_vulnerable_model=True, # model is already trained, so tolerate if vulnerable
sagemaker_session=sagemaker_session,
config_name=config_name,
)

# eula was already accepted if the model was successfully trained
Expand DownExpand Up@@ -1102,7 +1103,7 @@ def deploy(
tolerate_deprecated_model=self.tolerate_deprecated_model,
tolerate_vulnerable_model=self.tolerate_vulnerable_model,
sagemaker_session=self.sagemaker_session,
# config_name=self.config_name,
config_name=self.config_name,
)

# If a predictor class was passed, do not mutate predictor
Expand Down
5 changes: 4 additions & 1 deletion src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -478,7 +478,10 @@ def _add_tags_to_kwargs(kwargs: JumpStartEstimatorInitKwargs) -> JumpStartEstima

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version
kwargs.tags,
kwargs.model_id,
full_model_version,
config_name=kwargs.config_name,
)
return kwargs

Expand Down
2 changes: 1 addition & 1 deletion src/sagemaker/jumpstart/factory/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -496,7 +496,7 @@ def _add_tags_to_kwargs(kwargs: JumpStartModelDeployKwargs) -> Dict[str, Any]:

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type, kwargs.config_name
)

return kwargs
Expand Down
56 changes: 30 additions & 26 deletions src/sagemaker/jumpstart/session_utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -22,12 +22,12 @@
from sagemaker.utils import aws_partition


def get_model_id_version_from_endpoint(
def get_model_info_from_endpoint(
endpoint_name: str,
inference_component_name: Optional[str] = None,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str, Optional[str]]:
"""Given an endpoint and optionally inference component names, return the model IDand version.
) -> Tuple[str, str, Optional[str], Optional[str]]:
"""Optionally inference component names, return the model ID, version and config name.

Infers the model ID and version based on the resource tags. Returns a tuple of the model ID
and version. A third string element is included in the tuple for any inferred inference
Expand All@@ -46,30 +46,32 @@ def get_model_id_version_from_endpoint(
(
model_id,
model_version,
) = _get_model_id_version_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
config_name,
) = _get_model_info_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
inference_component_name, sagemaker_session
)

else:
(
model_id,
model_version,
config_name,
inference_component_name,
) = _get_model_id_version_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
) = _get_model_info_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
endpoint_name, sagemaker_session
)

else:
model_id, model_version = _get_model_id_version_from_model_based_endpoint(
model_id, model_version, config_name = _get_model_info_from_model_based_endpoint(
endpoint_name, inference_component_name, sagemaker_session
)
return model_id, model_version, inference_component_name
return model_id, model_version, inference_component_name, config_name


def _get_model_id_version_from_inference_component_endpoint_without_inference_component_name(
def _get_model_info_from_inference_component_endpoint_without_inference_component_name(
endpoint_name: str, sagemaker_session: Session
) -> Tuple[str, str, str]:
"""Given an endpoint name, derives the model ID, version, and inferred inference component name.
) -> Tuple[str, str, str, str]:
"""Derives the model ID, version, config name and inferred inference component name.

This function assumes the endpoint corresponds to an inference-component-based endpoint.
An endpoint is inference-component-based if and only if the associated endpoint config
Expand DownExpand Up@@ -98,14 +100,14 @@ def _get_model_id_version_from_inference_component_endpoint_without_inference_co
)
inference_component_name = inference_component_names[0]
return (
*_get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
*_get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name, sagemaker_session
),
inference_component_name,
)


def _get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
def _get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name: str, sagemaker_session: Session
):
"""Returns the model ID and version inferred from a SageMaker inference component.
Expand All@@ -123,7 +125,7 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
f"inference-component/{inference_component_name}"
)

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
inference_component_arn, sagemaker_session
)

Expand All@@ -134,15 +136,15 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
"when retrieving default predictor for this inference component."
)

return model_id, model_version
return model_id, model_version, config_name


def _get_model_id_version_from_model_based_endpoint(
def _get_model_info_from_model_based_endpoint(
endpoint_name: str,
inference_component_name: Optional[str],
sagemaker_session: Session,
) -> Tuple[str, str]:
"""Returns the model IDand version inferred from a model-based endpoint.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID, version and config name inferred from a model-based endpoint.

Raises:
ValueError: If an inference component name is supplied, or if the endpoint does
Expand All@@ -161,7 +163,7 @@ def _get_model_id_version_from_model_based_endpoint(

endpoint_arn = f"arn:{partition}:sagemaker:{region}:{account_id}:endpoint/{endpoint_name}"

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
endpoint_arn, sagemaker_session
)

Expand All@@ -172,14 +174,14 @@ def _get_model_id_version_from_model_based_endpoint(
"predictor for this endpoint."
)

return model_id, model_version
return model_id, model_version, config_name


def get_model_id_version_from_training_job(
def get_model_info_from_training_job(
training_job_name: str,
sagemaker_session: Optional[Session] = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str]:
"""Returns the model ID and version inferred from a training job.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID and version and config name inferred from a training job.

Raises:
ValueError: If the training job does not have tags from which the model ID
Expand All@@ -194,9 +196,11 @@ def get_model_id_version_from_training_job(
f"arn:{partition}:sagemaker:{region}:{account_id}:training-job/{training_job_name}"
)

model_id, inferred_model_version = get_jumpstart_model_id_version_from_resource_arn(
training_job_arn, sagemaker_session
)
(
model_id,
inferred_model_version,
config_name,
) = get_jumpstart_model_id_version_from_resource_arn(training_job_arn, sagemaker_session)

model_version = inferred_model_version or None

Expand All@@ -207,4 +211,4 @@ def get_model_id_version_from_training_job(
"for this training job."
)

return model_id, model_version
return model_id, model_version, config_name
22 changes: 12 additions & 10 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1064,9 +1064,8 @@ def from_json(self, json_obj: Dict[str, Any]) -> None:
Dictionary representation of the config component.
"""
for field in json_obj.keys():
if field not in self.__slots__:
raise ValueError(f"Invalid component field: {field}")
setattr(self, field, json_obj[field])
if field in self.__slots__:
setattr(self, field, json_obj[field])


class JumpStartMetadataConfig(JumpStartDataHolderType):
Expand DownExpand Up@@ -1164,20 +1163,17 @@ def get_top_config_from_ranking(
) -> Optional[JumpStartMetadataConfig]:
"""Gets the best the config based on config ranking.

Fallback to use the ordering in config names if
ranking is not available.
Args:
ranking_name (str):
The ranking name that config priority is based on.
instance_type (Optional[str]):
The instance type which the config selection is based on.

Raises:
ValueError: If the config exists but missing config ranking.
NotImplementedError: If the scope is unrecognized.
"""
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
raise ValueError(f"Config exists but missing config ranking {ranking_name}.")

if self.scope == JumpStartScriptScope.INFERENCE:
instance_type_attribute = "supported_inference_instance_types"
Expand All@@ -1186,8 +1182,14 @@ def get_top_config_from_ranking(
else:
raise NotImplementedError(f"Unknown script scope {self.scope}")

rankings = self.config_rankings.get(ranking_name)
for config_name in rankings.rankings:
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
ranked_config_names = sorted(list(self.configs.keys()))
else:
rankings = self.config_rankings.get(ranking_name)
ranked_config_names = rankings.rankings
for config_name in ranked_config_names:
resolved_config = self.configs[config_name].resolved_config
if instance_type and instance_type not in getattr(
resolved_config, instance_type_attribute
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
Merged
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
1 change: 1 addition & 0 deletions src/sagemaker/jumpstart/enums.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -92,6 +92,7 @@ class JumpStartTag(str, Enum):
MODEL_ID = "sagemaker-sdk:jumpstart-model-id"
MODEL_VERSION = "sagemaker-sdk:jumpstart-model-version"
MODEL_TYPE = "sagemaker-sdk:jumpstart-model-type"
MODEL_CONFIG_NAME = "sagemaker-sdk:jumpstart-model-config-name"


class SerializerType(str, Enum):
Expand Down
9 changes: 5 additions & 4 deletions src/sagemaker/jumpstart/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -33,7 +33,7 @@

from sagemaker.jumpstart.factory.estimator import get_deploy_kwargs, get_fit_kwargs, get_init_kwargs
from sagemaker.jumpstart.factory.model import get_default_predictor
from sagemaker.jumpstart.session_utils import get_model_id_version_from_training_job
from sagemaker.jumpstart.session_utils import get_model_info_from_training_job
from sagemaker.jumpstart.types import JumpStartMetadataConfig
from sagemaker.jumpstart.utils import (
get_jumpstart_configs,
Expand DownExpand Up@@ -730,10 +730,10 @@ def attach(
ValueError: if the model ID or version cannot be inferred from the training job.

"""

config_name = None
if model_id is None:

model_id, model_version = get_model_id_version_from_training_job(
model_id, model_version, config_name = get_model_info_from_training_job(
training_job_name=training_job_name, sagemaker_session=sagemaker_session
)

Expand All@@ -749,6 +749,7 @@ def attach(
tolerate_deprecated_model=True, # model is already trained, so tolerate if deprecated
tolerate_vulnerable_model=True, # model is already trained, so tolerate if vulnerable
sagemaker_session=sagemaker_session,
config_name=config_name,
)

# eula was already accepted if the model was successfully trained
Expand DownExpand Up@@ -1102,7 +1103,7 @@ def deploy(
tolerate_deprecated_model=self.tolerate_deprecated_model,
tolerate_vulnerable_model=self.tolerate_vulnerable_model,
sagemaker_session=self.sagemaker_session,
# config_name=self.config_name,
config_name=self.config_name,
)

# If a predictor class was passed, do not mutate predictor
Expand Down
5 changes: 4 additions & 1 deletion src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -478,7 +478,10 @@ def _add_tags_to_kwargs(kwargs: JumpStartEstimatorInitKwargs) -> JumpStartEstima

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version
kwargs.tags,
kwargs.model_id,
full_model_version,
config_name=kwargs.config_name,
)
return kwargs

Expand Down
2 changes: 1 addition & 1 deletion src/sagemaker/jumpstart/factory/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -496,7 +496,7 @@ def _add_tags_to_kwargs(kwargs: JumpStartModelDeployKwargs) -> Dict[str, Any]:

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type, kwargs.config_name
)

return kwargs
Expand Down
56 changes: 30 additions & 26 deletions src/sagemaker/jumpstart/session_utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -22,12 +22,12 @@
from sagemaker.utils import aws_partition


def get_model_id_version_from_endpoint(
def get_model_info_from_endpoint(
endpoint_name: str,
inference_component_name: Optional[str] = None,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str, Optional[str]]:
"""Given an endpoint and optionally inference component names, return the model IDand version.
) -> Tuple[str, str, Optional[str], Optional[str]]:
"""Optionally inference component names, return the model ID, version and config name.

Infers the model ID and version based on the resource tags. Returns a tuple of the model ID
and version. A third string element is included in the tuple for any inferred inference
Expand All@@ -46,30 +46,32 @@ def get_model_id_version_from_endpoint(
(
model_id,
model_version,
) = _get_model_id_version_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
config_name,
) = _get_model_info_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
inference_component_name, sagemaker_session
)

else:
(
model_id,
model_version,
config_name,
inference_component_name,
) = _get_model_id_version_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
) = _get_model_info_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
endpoint_name, sagemaker_session
)

else:
model_id, model_version = _get_model_id_version_from_model_based_endpoint(
model_id, model_version, config_name = _get_model_info_from_model_based_endpoint(
endpoint_name, inference_component_name, sagemaker_session
)
return model_id, model_version, inference_component_name
return model_id, model_version, inference_component_name, config_name


def _get_model_id_version_from_inference_component_endpoint_without_inference_component_name(
def _get_model_info_from_inference_component_endpoint_without_inference_component_name(
endpoint_name: str, sagemaker_session: Session
) -> Tuple[str, str, str]:
"""Given an endpoint name, derives the model ID, version, and inferred inference component name.
) -> Tuple[str, str, str, str]:
"""Derives the model ID, version, config name and inferred inference component name.

This function assumes the endpoint corresponds to an inference-component-based endpoint.
An endpoint is inference-component-based if and only if the associated endpoint config
Expand DownExpand Up@@ -98,14 +100,14 @@ def _get_model_id_version_from_inference_component_endpoint_without_inference_co
)
inference_component_name = inference_component_names[0]
return (
*_get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
*_get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name, sagemaker_session
),
inference_component_name,
)


def _get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
def _get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name: str, sagemaker_session: Session
):
"""Returns the model ID and version inferred from a SageMaker inference component.
Expand All@@ -123,7 +125,7 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
f"inference-component/{inference_component_name}"
)

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
inference_component_arn, sagemaker_session
)

Expand All@@ -134,15 +136,15 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
"when retrieving default predictor for this inference component."
)

return model_id, model_version
return model_id, model_version, config_name


def _get_model_id_version_from_model_based_endpoint(
def _get_model_info_from_model_based_endpoint(
endpoint_name: str,
inference_component_name: Optional[str],
sagemaker_session: Session,
) -> Tuple[str, str]:
"""Returns the model IDand version inferred from a model-based endpoint.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID, version and config name inferred from a model-based endpoint.

Raises:
ValueError: If an inference component name is supplied, or if the endpoint does
Expand All@@ -161,7 +163,7 @@ def _get_model_id_version_from_model_based_endpoint(

endpoint_arn = f"arn:{partition}:sagemaker:{region}:{account_id}:endpoint/{endpoint_name}"

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
endpoint_arn, sagemaker_session
)

Expand All@@ -172,14 +174,14 @@ def _get_model_id_version_from_model_based_endpoint(
"predictor for this endpoint."
)

return model_id, model_version
return model_id, model_version, config_name


def get_model_id_version_from_training_job(
def get_model_info_from_training_job(
training_job_name: str,
sagemaker_session: Optional[Session] = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str]:
"""Returns the model ID and version inferred from a training job.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID and version and config name inferred from a training job.

Raises:
ValueError: If the training job does not have tags from which the model ID
Expand All@@ -194,9 +196,11 @@ def get_model_id_version_from_training_job(
f"arn:{partition}:sagemaker:{region}:{account_id}:training-job/{training_job_name}"
)

model_id, inferred_model_version = get_jumpstart_model_id_version_from_resource_arn(
training_job_arn, sagemaker_session
)
(
model_id,
inferred_model_version,
config_name,
) = get_jumpstart_model_id_version_from_resource_arn(training_job_arn, sagemaker_session)

model_version = inferred_model_version or None

Expand All@@ -207,4 +211,4 @@ def get_model_id_version_from_training_job(
"for this training job."
)

return model_id, model_version
return model_id, model_version, config_name
22 changes: 12 additions & 10 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1064,9 +1064,8 @@ def from_json(self, json_obj: Dict[str, Any]) -> None:
Dictionary representation of the config component.
"""
for field in json_obj.keys():
if field not in self.__slots__:
raise ValueError(f"Invalid component field: {field}")
setattr(self, field, json_obj[field])
if field in self.__slots__:
setattr(self, field, json_obj[field])


class JumpStartMetadataConfig(JumpStartDataHolderType):
Expand DownExpand Up@@ -1164,20 +1163,17 @@ def get_top_config_from_ranking(
) -> Optional[JumpStartMetadataConfig]:
"""Gets the best the config based on config ranking.

Fallback to use the ordering in config names if
ranking is not available.
Args:
ranking_name (str):
The ranking name that config priority is based on.
instance_type (Optional[str]):
The instance type which the config selection is based on.

Raises:
ValueError: If the config exists but missing config ranking.
NotImplementedError: If the scope is unrecognized.
"""
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
raise ValueError(f"Config exists but missing config ranking {ranking_name}.")

if self.scope == JumpStartScriptScope.INFERENCE:
instance_type_attribute = "supported_inference_instance_types"
Expand All@@ -1186,8 +1182,14 @@ def get_top_config_from_ranking(
else:
raise NotImplementedError(f"Unknown script scope {self.scope}")

rankings = self.config_rankings.get(ranking_name)
for config_name in rankings.rankings:
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
ranked_config_names = sorted(list(self.configs.keys()))
else:
rankings = self.config_rankings.get(ranking_name)
ranked_config_names = rankings.rankings
for config_name in ranked_config_names:
resolved_config = self.configs[config_name].resolved_config
if instance_type and instance_type not in getattr(
resolved_config, instance_type_attribute
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
Merged
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
1 change: 1 addition & 0 deletions src/sagemaker/jumpstart/enums.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -92,6 +92,7 @@ class JumpStartTag(str, Enum):
MODEL_ID = "sagemaker-sdk:jumpstart-model-id"
MODEL_VERSION = "sagemaker-sdk:jumpstart-model-version"
MODEL_TYPE = "sagemaker-sdk:jumpstart-model-type"
MODEL_CONFIG_NAME = "sagemaker-sdk:jumpstart-model-config-name"


class SerializerType(str, Enum):
Expand Down
9 changes: 5 additions & 4 deletions src/sagemaker/jumpstart/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -33,7 +33,7 @@

from sagemaker.jumpstart.factory.estimator import get_deploy_kwargs, get_fit_kwargs, get_init_kwargs
from sagemaker.jumpstart.factory.model import get_default_predictor
from sagemaker.jumpstart.session_utils import get_model_id_version_from_training_job
from sagemaker.jumpstart.session_utils import get_model_info_from_training_job
from sagemaker.jumpstart.types import JumpStartMetadataConfig
from sagemaker.jumpstart.utils import (
get_jumpstart_configs,
Expand DownExpand Up@@ -730,10 +730,10 @@ def attach(
ValueError: if the model ID or version cannot be inferred from the training job.

"""

config_name = None
if model_id is None:

model_id, model_version = get_model_id_version_from_training_job(
model_id, model_version, config_name = get_model_info_from_training_job(
training_job_name=training_job_name, sagemaker_session=sagemaker_session
)

Expand All@@ -749,6 +749,7 @@ def attach(
tolerate_deprecated_model=True, # model is already trained, so tolerate if deprecated
tolerate_vulnerable_model=True, # model is already trained, so tolerate if vulnerable
sagemaker_session=sagemaker_session,
config_name=config_name,
)

# eula was already accepted if the model was successfully trained
Expand DownExpand Up@@ -1102,7 +1103,7 @@ def deploy(
tolerate_deprecated_model=self.tolerate_deprecated_model,
tolerate_vulnerable_model=self.tolerate_vulnerable_model,
sagemaker_session=self.sagemaker_session,
# config_name=self.config_name,
config_name=self.config_name,
)

# If a predictor class was passed, do not mutate predictor
Expand Down
5 changes: 4 additions & 1 deletion src/sagemaker/jumpstart/factory/estimator.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -478,7 +478,10 @@ def _add_tags_to_kwargs(kwargs: JumpStartEstimatorInitKwargs) -> JumpStartEstima

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version
kwargs.tags,
kwargs.model_id,
full_model_version,
config_name=kwargs.config_name,
)
return kwargs

Expand Down
2 changes: 1 addition & 1 deletion src/sagemaker/jumpstart/factory/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -496,7 +496,7 @@ def _add_tags_to_kwargs(kwargs: JumpStartModelDeployKwargs) -> Dict[str, Any]:

if kwargs.sagemaker_session.settings.include_jumpstart_tags:
kwargs.tags = add_jumpstart_model_id_version_tags(
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type
kwargs.tags, kwargs.model_id, full_model_version, kwargs.model_type, kwargs.config_name
)

return kwargs
Expand Down
56 changes: 30 additions & 26 deletions src/sagemaker/jumpstart/session_utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -22,12 +22,12 @@
from sagemaker.utils import aws_partition


def get_model_id_version_from_endpoint(
def get_model_info_from_endpoint(
endpoint_name: str,
inference_component_name: Optional[str] = None,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str, Optional[str]]:
"""Given an endpoint and optionally inference component names, return the model IDand version.
) -> Tuple[str, str, Optional[str], Optional[str]]:
"""Optionally inference component names, return the model ID, version and config name.

Infers the model ID and version based on the resource tags. Returns a tuple of the model ID
and version. A third string element is included in the tuple for any inferred inference
Expand All@@ -46,30 +46,32 @@ def get_model_id_version_from_endpoint(
(
model_id,
model_version,
) = _get_model_id_version_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
config_name,
) = _get_model_info_from_inference_component_endpoint_with_inference_component_name( # noqa E501 # pylint: disable=c0301
inference_component_name, sagemaker_session
)

else:
(
model_id,
model_version,
config_name,
inference_component_name,
) = _get_model_id_version_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
) = _get_model_info_from_inference_component_endpoint_without_inference_component_name( # noqa E501 # pylint: disable=c0301
endpoint_name, sagemaker_session
)

else:
model_id, model_version = _get_model_id_version_from_model_based_endpoint(
model_id, model_version, config_name = _get_model_info_from_model_based_endpoint(
endpoint_name, inference_component_name, sagemaker_session
)
return model_id, model_version, inference_component_name
return model_id, model_version, inference_component_name, config_name


def _get_model_id_version_from_inference_component_endpoint_without_inference_component_name(
def _get_model_info_from_inference_component_endpoint_without_inference_component_name(
endpoint_name: str, sagemaker_session: Session
) -> Tuple[str, str, str]:
"""Given an endpoint name, derives the model ID, version, and inferred inference component name.
) -> Tuple[str, str, str, str]:
"""Derives the model ID, version, config name and inferred inference component name.

This function assumes the endpoint corresponds to an inference-component-based endpoint.
An endpoint is inference-component-based if and only if the associated endpoint config
Expand DownExpand Up@@ -98,14 +100,14 @@ def _get_model_id_version_from_inference_component_endpoint_without_inference_co
)
inference_component_name = inference_component_names[0]
return (
*_get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
*_get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name, sagemaker_session
),
inference_component_name,
)


def _get_model_id_version_from_inference_component_endpoint_with_inference_component_name(
def _get_model_info_from_inference_component_endpoint_with_inference_component_name(
inference_component_name: str, sagemaker_session: Session
):
"""Returns the model ID and version inferred from a SageMaker inference component.
Expand All@@ -123,7 +125,7 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
f"inference-component/{inference_component_name}"
)

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
inference_component_arn, sagemaker_session
)

Expand All@@ -134,15 +136,15 @@ def _get_model_id_version_from_inference_component_endpoint_with_inference_compo
"when retrieving default predictor for this inference component."
)

return model_id, model_version
return model_id, model_version, config_name


def _get_model_id_version_from_model_based_endpoint(
def _get_model_info_from_model_based_endpoint(
endpoint_name: str,
inference_component_name: Optional[str],
sagemaker_session: Session,
) -> Tuple[str, str]:
"""Returns the model IDand version inferred from a model-based endpoint.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID, version and config name inferred from a model-based endpoint.

Raises:
ValueError: If an inference component name is supplied, or if the endpoint does
Expand All@@ -161,7 +163,7 @@ def _get_model_id_version_from_model_based_endpoint(

endpoint_arn = f"arn:{partition}:sagemaker:{region}:{account_id}:endpoint/{endpoint_name}"

model_id, model_version = get_jumpstart_model_id_version_from_resource_arn(
model_id, model_version, config_name = get_jumpstart_model_id_version_from_resource_arn(
endpoint_arn, sagemaker_session
)

Expand All@@ -172,14 +174,14 @@ def _get_model_id_version_from_model_based_endpoint(
"predictor for this endpoint."
)

return model_id, model_version
return model_id, model_version, config_name


def get_model_id_version_from_training_job(
def get_model_info_from_training_job(
training_job_name: str,
sagemaker_session: Optional[Session] = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
) -> Tuple[str, str]:
"""Returns the model ID and version inferred from a training job.
) -> Tuple[str, str, Optional[str]]:
"""Returns the model ID and version and config name inferred from a training job.

Raises:
ValueError: If the training job does not have tags from which the model ID
Expand All@@ -194,9 +196,11 @@ def get_model_id_version_from_training_job(
f"arn:{partition}:sagemaker:{region}:{account_id}:training-job/{training_job_name}"
)

model_id, inferred_model_version = get_jumpstart_model_id_version_from_resource_arn(
training_job_arn, sagemaker_session
)
(
model_id,
inferred_model_version,
config_name,
) = get_jumpstart_model_id_version_from_resource_arn(training_job_arn, sagemaker_session)

model_version = inferred_model_version or None

Expand All@@ -207,4 +211,4 @@ def get_model_id_version_from_training_job(
"for this training job."
)

return model_id, model_version
return model_id, model_version, config_name
22 changes: 12 additions & 10 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1064,9 +1064,8 @@ def from_json(self, json_obj: Dict[str, Any]) -> None:
Dictionary representation of the config component.
"""
for field in json_obj.keys():
if field not in self.__slots__:
raise ValueError(f"Invalid component field: {field}")
setattr(self, field, json_obj[field])
if field in self.__slots__:
setattr(self, field, json_obj[field])


class JumpStartMetadataConfig(JumpStartDataHolderType):
Expand DownExpand Up@@ -1164,20 +1163,17 @@ def get_top_config_from_ranking(
) -> Optional[JumpStartMetadataConfig]:
"""Gets the best the config based on config ranking.

Fallback to use the ordering in config names if
ranking is not available.
Args:
ranking_name (str):
The ranking name that config priority is based on.
instance_type (Optional[str]):
The instance type which the config selection is based on.

Raises:
ValueError: If the config exists but missing config ranking.
NotImplementedError: If the scope is unrecognized.
"""
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
raise ValueError(f"Config exists but missing config ranking {ranking_name}.")

if self.scope == JumpStartScriptScope.INFERENCE:
instance_type_attribute = "supported_inference_instance_types"
Expand All@@ -1186,8 +1182,14 @@ def get_top_config_from_ranking(
else:
raise NotImplementedError(f"Unknown script scope {self.scope}")

rankings = self.config_rankings.get(ranking_name)
for config_name in rankings.rankings:
if self.configs and (
not self.config_rankings or not self.config_rankings.get(ranking_name)
):
ranked_config_names = sorted(list(self.configs.keys()))
else:
rankings = self.config_rankings.get(ranking_name)
ranked_config_names = rankings.rankings
for config_name in ranked_config_names:
resolved_config = self.configs[config_name].resolved_config
if instance_type and instance_type not in getattr(
resolved_config, instance_type_attribute
Expand Down
Loading