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
109 changes: 106 additions & 3 deletions src/sagemaker/jumpstart/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,7 +14,9 @@

from __future__ import absolute_import

from typing import Dict, List, Optional, Union
from functools import lru_cache
from typing import Dict, List, Optional, Union, Any
import pandas as pd
from botocore.exceptions import ClientError

from sagemaker import payloads
Expand All@@ -36,14 +38,21 @@
get_init_kwargs,
get_register_kwargs,
)
from sagemaker.jumpstart.types import JumpStartSerializablePayload
from sagemaker.jumpstart.types import (
JumpStartSerializablePayload,
DeploymentConfigMetadata,
JumpStartBenchmarkStat,
JumpStartMetadataConfig,
)
from sagemaker.jumpstart.utils import (
validate_model_id_and_get_type,
verify_model_region_and_return_specs,
get_jumpstart_configs,
extract_metrics_from_deployment_configs,
)
from sagemaker.jumpstart.constants import JUMPSTART_LOGGER
from sagemaker.jumpstart.enums import JumpStartModelType
from sagemaker.utils import stringify_object, format_tags, Tags
from sagemaker.utils import stringify_object, format_tags, Tags, get_instance_rate_per_hour
from sagemaker.model import (
Model,
ModelPackage,
Expand DownExpand Up@@ -352,6 +361,18 @@ def _validate_model_id_and_type():
self.model_package_arn = model_init_kwargs.model_package_arn
self.init_kwargs = model_init_kwargs.to_kwargs_dict(False)

metadata_configs = get_jumpstart_configs(
region=self.region,
model_id=self.model_id,
model_version=self.model_version,
sagemaker_session=self.sagemaker_session,
model_type=self.model_type,
)
self._deployment_configs = [
self._convert_to_deployment_config_metadata(config_name, config)
for config_name, config in metadata_configs.items()
]

def log_subscription_warning(self) -> None:
"""Log message prompting the customer to subscribe to the proprietary model."""
subscription_link = verify_model_region_and_return_specs(
Expand DownExpand Up@@ -420,6 +441,27 @@ def set_deployment_config(self, config_name: Optional[str]) -> None:
model_id=self.model_id, model_version=self.model_version, config_name=config_name
)

@property
def benchmark_metrics(self) -> pd.DataFrame:
"""Benchmark Metrics for deployment configs

Returns:
Metrics: Pandas DataFrame object.
"""
return pd.DataFrame(self._get_benchmark_data(self.config_name))

def display_benchmark_metrics(self) -> None:
"""Display Benchmark Metrics for deployment configs."""
print(self.benchmark_metrics.to_markdown())

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self._deployment_configs

def _create_sagemaker_model(
self,
instance_type=None,
Expand DownExpand Up@@ -808,6 +850,67 @@ def register_deploy_wrapper(*args, **kwargs):

return model_package

@lru_cache
def _get_benchmark_data(self, config_name: str) -> Dict[str, List[str]]:
"""Constructs deployment configs benchmark data.

Args:
config_name (str): The name of the selected deployment config.
Returns:
Dict[str, List[str]]: Deployment config benchmark data.
"""
return extract_metrics_from_deployment_configs(
self._deployment_configs,
config_name,
)

def _convert_to_deployment_config_metadata(
self, config_name: str, metadata_config: JumpStartMetadataConfig
) -> Dict[str, Any]:
"""Retrieve deployment config for config name.

Args:
config_name (str): Name of deployment config.
metadata_config (JumpStartMetadataConfig): Metadata config for deployment config.
Returns:
A deployment metadata config for config name (dict[str, Any]).
"""
default_inference_instance_type = metadata_config.resolved_config.get(
"default_inference_instance_type"
)

instance_rate = get_instance_rate_per_hour(
instance_type=default_inference_instance_type, region=self.region
)

benchmark_metrics = (
metadata_config.benchmark_metrics.get(default_inference_instance_type)
if metadata_config.benchmark_metrics is not None
else None
)
if instance_rate is not None:
if benchmark_metrics is not None:
benchmark_metrics.append(JumpStartBenchmarkStat(instance_rate))
else:
benchmark_metrics = [JumpStartBenchmarkStat(instance_rate)]

init_kwargs = get_init_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)
deploy_kwargs = get_deploy_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)

deployment_config_metadata = DeploymentConfigMetadata(
config_name, benchmark_metrics, init_kwargs, deploy_kwargs
)

return deployment_config_metadata.to_json()

def __str__(self) -> str:
"""Overriding str(*) method to make more human-readable."""
return stringify_object(self)
96 changes: 96 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2206,3 +2206,99 @@ def __init__(
self.skip_model_validation = skip_model_validation
self.source_uri = source_uri
self.config_name = config_name


class BaseDeploymentConfigDataHolder(JumpStartDataHolderType):
"""Base class for Deployment Config Data."""

def _convert_to_pascal_case(self, attr_name: str) -> str:
"""Converts a snake_case attribute name into a camelCased string.

Args:
attr_name (str): The snake_case attribute name.
Returns:
str: The PascalCased attribute name.
"""
return attr_name.replace("_", " ").title().replace(" ", "")

def to_json(self) -> Dict[str, Any]:
"""Represents ``This`` object as JSON."""
json_obj = {}
for att in self.__slots__:
if hasattr(self, att):
cur_val = getattr(self, att)
att = self._convert_to_pascal_case(att)
if issubclass(type(cur_val), JumpStartDataHolderType):
json_obj[att] = cur_val.to_json()
elif isinstance(cur_val, list):
json_obj[att] = []
for obj in cur_val:
if issubclass(type(obj), JumpStartDataHolderType):

@evakravievakraviApr 23, 2024

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

this logic's really complicated. if you can find a way to reduce indentation level, that'd improve readability

json_obj[att].append(obj.to_json())
else:
json_obj[att].append(obj)
elif isinstance(cur_val, dict):
json_obj[att] = {}
for key, val in cur_val.items():
if issubclass(type(val), JumpStartDataHolderType):
json_obj[att][self._convert_to_pascal_case(key)] = val.to_json()
else:
json_obj[att][key] = val
else:
json_obj[att] = cur_val
return json_obj


class DeploymentConfig(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config."""

__slots__ = [
"model_data_download_timeout",
"container_startup_health_check_timeout",
"image_uri",
"model_data",
"instance_type",
"environment",
"compute_resource_requirements",
]

def __init__(
self, init_kwargs: JumpStartModelInitKwargs, deploy_kwargs: JumpStartModelDeployKwargs
):
"""Instantiates DeploymentConfig object."""
if init_kwargs is not None:
self.image_uri = init_kwargs.image_uri
self.model_data = init_kwargs.model_data
self.instance_type = init_kwargs.instance_type
self.environment = init_kwargs.env
if init_kwargs.resources is not None:
self.compute_resource_requirements = (
init_kwargs.resources.get_compute_resource_requirements()
)
if deploy_kwargs is not None:
self.model_data_download_timeout = deploy_kwargs.model_data_download_timeout
self.container_startup_health_check_timeout = (
deploy_kwargs.container_startup_health_check_timeout
)


class DeploymentConfigMetadata(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config Metadata"""

__slots__ = [
"config_name",
"benchmark_metrics",
"deployment_config",
]

def __init__(
self,
config_name: str,
benchmark_metrics: List[JumpStartBenchmarkStat],
init_kwargs: JumpStartModelInitKwargs,
deploy_kwargs: JumpStartModelDeployKwargs,
):
"""Instantiates DeploymentConfigMetadata object."""
self.config_name = config_name
self.benchmark_metrics = benchmark_metrics
self.deployment_config = DeploymentConfig(init_kwargs, deploy_kwargs)
41 changes: 41 additions & 0 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -999,3 +999,44 @@ def get_jumpstart_configs(
if metadata_configs
else {}
)


def extract_metrics_from_deployment_configs(
deployment_configs: List[Dict[str, Any]], config_name: str
) -> Dict[str, List[str]]:
"""Extracts metrics from deployment configs.

Args:
deployment_configs (list[dict[str, Any]]): List of deployment configs.
config_name (str): The name of the deployment config use by the model.
"""

data = {"Config Name": [], "Instance Type": [], "Selected": []}

for index, deployment_config in enumerate(deployment_configs):
if deployment_config.get("DeploymentConfig") is None:
continue

benchmark_metrics = deployment_config.get("BenchmarkMetrics")
if benchmark_metrics is not None:
data["Config Name"].append(deployment_config.get("ConfigName"))
data["Instance Type"].append(
deployment_config.get("DeploymentConfig").get("InstanceType")
Comment thread
makungaj1 marked this conversation as resolved.
)
data["Selected"].append(
"Yes"
if (config_name is not None and config_name == deployment_config.get("ConfigName"))
else "No"
)

if index == 0:
for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
data[column_name] = []

for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
if column_name in data.keys():
data[column_name].append(benchmark_metric.get("value"))

return data
14 changes: 13 additions & 1 deletion src/sagemaker/serve/builder/jumpstart_builder.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,7 +16,7 @@
import copy
from abc import ABC, abstractmethod
from datetime import datetime, timedelta
from typing import Type
from typing import Type, Any, List, Dict
import logging

from sagemaker.model import Model
Expand DownExpand Up@@ -431,6 +431,18 @@ def tune_for_tgi_jumpstart(self, max_tuning_duration: int = 1800):
sharded_supported=sharded_supported, max_tuning_duration=max_tuning_duration
)

def display_benchmark_metrics(self):
"""Display Markdown Benchmark Metrics for deployment configs."""
self.pysdk_model.display_benchmark_metrics()

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model in the current region.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self.pysdk_model.list_deployment_configs()

def _build_for_jumpstart(self):
"""Placeholder docstring"""
# we do not pickle for jumpstart. set to none
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
109 changes: 106 additions & 3 deletions src/sagemaker/jumpstart/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,7 +14,9 @@

from __future__ import absolute_import

from typing import Dict, List, Optional, Union
from functools import lru_cache
from typing import Dict, List, Optional, Union, Any
import pandas as pd
from botocore.exceptions import ClientError

from sagemaker import payloads
Expand All@@ -36,14 +38,21 @@
get_init_kwargs,
get_register_kwargs,
)
from sagemaker.jumpstart.types import JumpStartSerializablePayload
from sagemaker.jumpstart.types import (
JumpStartSerializablePayload,
DeploymentConfigMetadata,
JumpStartBenchmarkStat,
JumpStartMetadataConfig,
)
from sagemaker.jumpstart.utils import (
validate_model_id_and_get_type,
verify_model_region_and_return_specs,
get_jumpstart_configs,
extract_metrics_from_deployment_configs,
)
from sagemaker.jumpstart.constants import JUMPSTART_LOGGER
from sagemaker.jumpstart.enums import JumpStartModelType
from sagemaker.utils import stringify_object, format_tags, Tags
from sagemaker.utils import stringify_object, format_tags, Tags, get_instance_rate_per_hour
from sagemaker.model import (
Model,
ModelPackage,
Expand DownExpand Up@@ -352,6 +361,18 @@ def _validate_model_id_and_type():
self.model_package_arn = model_init_kwargs.model_package_arn
self.init_kwargs = model_init_kwargs.to_kwargs_dict(False)

metadata_configs = get_jumpstart_configs(
region=self.region,
model_id=self.model_id,
model_version=self.model_version,
sagemaker_session=self.sagemaker_session,
model_type=self.model_type,
)
self._deployment_configs = [
self._convert_to_deployment_config_metadata(config_name, config)
for config_name, config in metadata_configs.items()
]

def log_subscription_warning(self) -> None:
"""Log message prompting the customer to subscribe to the proprietary model."""
subscription_link = verify_model_region_and_return_specs(
Expand DownExpand Up@@ -420,6 +441,27 @@ def set_deployment_config(self, config_name: Optional[str]) -> None:
model_id=self.model_id, model_version=self.model_version, config_name=config_name
)

@property
def benchmark_metrics(self) -> pd.DataFrame:
"""Benchmark Metrics for deployment configs

Returns:
Metrics: Pandas DataFrame object.
"""
return pd.DataFrame(self._get_benchmark_data(self.config_name))

def display_benchmark_metrics(self) -> None:
"""Display Benchmark Metrics for deployment configs."""
print(self.benchmark_metrics.to_markdown())

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self._deployment_configs

def _create_sagemaker_model(
self,
instance_type=None,
Expand DownExpand Up@@ -808,6 +850,67 @@ def register_deploy_wrapper(*args, **kwargs):

return model_package

@lru_cache
def _get_benchmark_data(self, config_name: str) -> Dict[str, List[str]]:
"""Constructs deployment configs benchmark data.

Args:
config_name (str): The name of the selected deployment config.
Returns:
Dict[str, List[str]]: Deployment config benchmark data.
"""
return extract_metrics_from_deployment_configs(
self._deployment_configs,
config_name,
)

def _convert_to_deployment_config_metadata(
self, config_name: str, metadata_config: JumpStartMetadataConfig
) -> Dict[str, Any]:
"""Retrieve deployment config for config name.

Args:
config_name (str): Name of deployment config.
metadata_config (JumpStartMetadataConfig): Metadata config for deployment config.
Returns:
A deployment metadata config for config name (dict[str, Any]).
"""
default_inference_instance_type = metadata_config.resolved_config.get(
"default_inference_instance_type"
)

instance_rate = get_instance_rate_per_hour(
instance_type=default_inference_instance_type, region=self.region
)

benchmark_metrics = (
metadata_config.benchmark_metrics.get(default_inference_instance_type)
if metadata_config.benchmark_metrics is not None
else None
)
if instance_rate is not None:
if benchmark_metrics is not None:
benchmark_metrics.append(JumpStartBenchmarkStat(instance_rate))
else:
benchmark_metrics = [JumpStartBenchmarkStat(instance_rate)]

init_kwargs = get_init_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)
deploy_kwargs = get_deploy_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)

deployment_config_metadata = DeploymentConfigMetadata(
config_name, benchmark_metrics, init_kwargs, deploy_kwargs
)

return deployment_config_metadata.to_json()

def __str__(self) -> str:
"""Overriding str(*) method to make more human-readable."""
return stringify_object(self)
96 changes: 96 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2206,3 +2206,99 @@ def __init__(
self.skip_model_validation = skip_model_validation
self.source_uri = source_uri
self.config_name = config_name


class BaseDeploymentConfigDataHolder(JumpStartDataHolderType):
"""Base class for Deployment Config Data."""

def _convert_to_pascal_case(self, attr_name: str) -> str:
"""Converts a snake_case attribute name into a camelCased string.

Args:
attr_name (str): The snake_case attribute name.
Returns:
str: The PascalCased attribute name.
"""
return attr_name.replace("_", " ").title().replace(" ", "")

def to_json(self) -> Dict[str, Any]:
"""Represents ``This`` object as JSON."""
json_obj = {}
for att in self.__slots__:
if hasattr(self, att):
cur_val = getattr(self, att)
att = self._convert_to_pascal_case(att)
if issubclass(type(cur_val), JumpStartDataHolderType):
json_obj[att] = cur_val.to_json()
elif isinstance(cur_val, list):
json_obj[att] = []
for obj in cur_val:
if issubclass(type(obj), JumpStartDataHolderType):

@evakravievakraviApr 23, 2024

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

this logic's really complicated. if you can find a way to reduce indentation level, that'd improve readability

json_obj[att].append(obj.to_json())
else:
json_obj[att].append(obj)
elif isinstance(cur_val, dict):
json_obj[att] = {}
for key, val in cur_val.items():
if issubclass(type(val), JumpStartDataHolderType):
json_obj[att][self._convert_to_pascal_case(key)] = val.to_json()
else:
json_obj[att][key] = val
else:
json_obj[att] = cur_val
return json_obj


class DeploymentConfig(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config."""

__slots__ = [
"model_data_download_timeout",
"container_startup_health_check_timeout",
"image_uri",
"model_data",
"instance_type",
"environment",
"compute_resource_requirements",
]

def __init__(
self, init_kwargs: JumpStartModelInitKwargs, deploy_kwargs: JumpStartModelDeployKwargs
):
"""Instantiates DeploymentConfig object."""
if init_kwargs is not None:
self.image_uri = init_kwargs.image_uri
self.model_data = init_kwargs.model_data
self.instance_type = init_kwargs.instance_type
self.environment = init_kwargs.env
if init_kwargs.resources is not None:
self.compute_resource_requirements = (
init_kwargs.resources.get_compute_resource_requirements()
)
if deploy_kwargs is not None:
self.model_data_download_timeout = deploy_kwargs.model_data_download_timeout
self.container_startup_health_check_timeout = (
deploy_kwargs.container_startup_health_check_timeout
)


class DeploymentConfigMetadata(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config Metadata"""

__slots__ = [
"config_name",
"benchmark_metrics",
"deployment_config",
]

def __init__(
self,
config_name: str,
benchmark_metrics: List[JumpStartBenchmarkStat],
init_kwargs: JumpStartModelInitKwargs,
deploy_kwargs: JumpStartModelDeployKwargs,
):
"""Instantiates DeploymentConfigMetadata object."""
self.config_name = config_name
self.benchmark_metrics = benchmark_metrics
self.deployment_config = DeploymentConfig(init_kwargs, deploy_kwargs)
41 changes: 41 additions & 0 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -999,3 +999,44 @@ def get_jumpstart_configs(
if metadata_configs
else {}
)


def extract_metrics_from_deployment_configs(
deployment_configs: List[Dict[str, Any]], config_name: str
) -> Dict[str, List[str]]:
"""Extracts metrics from deployment configs.

Args:
deployment_configs (list[dict[str, Any]]): List of deployment configs.
config_name (str): The name of the deployment config use by the model.
"""

data = {"Config Name": [], "Instance Type": [], "Selected": []}

for index, deployment_config in enumerate(deployment_configs):
if deployment_config.get("DeploymentConfig") is None:
continue

benchmark_metrics = deployment_config.get("BenchmarkMetrics")
if benchmark_metrics is not None:
data["Config Name"].append(deployment_config.get("ConfigName"))
data["Instance Type"].append(
deployment_config.get("DeploymentConfig").get("InstanceType")
Comment thread
makungaj1 marked this conversation as resolved.
)
data["Selected"].append(
"Yes"
if (config_name is not None and config_name == deployment_config.get("ConfigName"))
else "No"
)

if index == 0:
for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
data[column_name] = []

for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
if column_name in data.keys():
data[column_name].append(benchmark_metric.get("value"))

return data
14 changes: 13 additions & 1 deletion src/sagemaker/serve/builder/jumpstart_builder.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,7 +16,7 @@
import copy
from abc import ABC, abstractmethod
from datetime import datetime, timedelta
from typing import Type
from typing import Type, Any, List, Dict
import logging

from sagemaker.model import Model
Expand DownExpand Up@@ -431,6 +431,18 @@ def tune_for_tgi_jumpstart(self, max_tuning_duration: int = 1800):
sharded_supported=sharded_supported, max_tuning_duration=max_tuning_duration
)

def display_benchmark_metrics(self):
"""Display Markdown Benchmark Metrics for deployment configs."""
self.pysdk_model.display_benchmark_metrics()

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model in the current region.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self.pysdk_model.list_deployment_configs()

def _build_for_jumpstart(self):
"""Placeholder docstring"""
# we do not pickle for jumpstart. set to none
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
109 changes: 106 additions & 3 deletions src/sagemaker/jumpstart/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,7 +14,9 @@

from __future__ import absolute_import

from typing import Dict, List, Optional, Union
from functools import lru_cache
from typing import Dict, List, Optional, Union, Any
import pandas as pd
from botocore.exceptions import ClientError

from sagemaker import payloads
Expand All@@ -36,14 +38,21 @@
get_init_kwargs,
get_register_kwargs,
)
from sagemaker.jumpstart.types import JumpStartSerializablePayload
from sagemaker.jumpstart.types import (
JumpStartSerializablePayload,
DeploymentConfigMetadata,
JumpStartBenchmarkStat,
JumpStartMetadataConfig,
)
from sagemaker.jumpstart.utils import (
validate_model_id_and_get_type,
verify_model_region_and_return_specs,
get_jumpstart_configs,
extract_metrics_from_deployment_configs,
)
from sagemaker.jumpstart.constants import JUMPSTART_LOGGER
from sagemaker.jumpstart.enums import JumpStartModelType
from sagemaker.utils import stringify_object, format_tags, Tags
from sagemaker.utils import stringify_object, format_tags, Tags, get_instance_rate_per_hour
from sagemaker.model import (
Model,
ModelPackage,
Expand DownExpand Up@@ -352,6 +361,18 @@ def _validate_model_id_and_type():
self.model_package_arn = model_init_kwargs.model_package_arn
self.init_kwargs = model_init_kwargs.to_kwargs_dict(False)

metadata_configs = get_jumpstart_configs(
region=self.region,
model_id=self.model_id,
model_version=self.model_version,
sagemaker_session=self.sagemaker_session,
model_type=self.model_type,
)
self._deployment_configs = [
self._convert_to_deployment_config_metadata(config_name, config)
for config_name, config in metadata_configs.items()
]

def log_subscription_warning(self) -> None:
"""Log message prompting the customer to subscribe to the proprietary model."""
subscription_link = verify_model_region_and_return_specs(
Expand DownExpand Up@@ -420,6 +441,27 @@ def set_deployment_config(self, config_name: Optional[str]) -> None:
model_id=self.model_id, model_version=self.model_version, config_name=config_name
)

@property
def benchmark_metrics(self) -> pd.DataFrame:
"""Benchmark Metrics for deployment configs

Returns:
Metrics: Pandas DataFrame object.
"""
return pd.DataFrame(self._get_benchmark_data(self.config_name))

def display_benchmark_metrics(self) -> None:
"""Display Benchmark Metrics for deployment configs."""
print(self.benchmark_metrics.to_markdown())

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self._deployment_configs

def _create_sagemaker_model(
self,
instance_type=None,
Expand DownExpand Up@@ -808,6 +850,67 @@ def register_deploy_wrapper(*args, **kwargs):

return model_package

@lru_cache
def _get_benchmark_data(self, config_name: str) -> Dict[str, List[str]]:
"""Constructs deployment configs benchmark data.

Args:
config_name (str): The name of the selected deployment config.
Returns:
Dict[str, List[str]]: Deployment config benchmark data.
"""
return extract_metrics_from_deployment_configs(
self._deployment_configs,
config_name,
)

def _convert_to_deployment_config_metadata(
self, config_name: str, metadata_config: JumpStartMetadataConfig
) -> Dict[str, Any]:
"""Retrieve deployment config for config name.

Args:
config_name (str): Name of deployment config.
metadata_config (JumpStartMetadataConfig): Metadata config for deployment config.
Returns:
A deployment metadata config for config name (dict[str, Any]).
"""
default_inference_instance_type = metadata_config.resolved_config.get(
"default_inference_instance_type"
)

instance_rate = get_instance_rate_per_hour(
instance_type=default_inference_instance_type, region=self.region
)

benchmark_metrics = (
metadata_config.benchmark_metrics.get(default_inference_instance_type)
if metadata_config.benchmark_metrics is not None
else None
)
if instance_rate is not None:
if benchmark_metrics is not None:
benchmark_metrics.append(JumpStartBenchmarkStat(instance_rate))
else:
benchmark_metrics = [JumpStartBenchmarkStat(instance_rate)]

init_kwargs = get_init_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)
deploy_kwargs = get_deploy_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)

deployment_config_metadata = DeploymentConfigMetadata(
config_name, benchmark_metrics, init_kwargs, deploy_kwargs
)

return deployment_config_metadata.to_json()

def __str__(self) -> str:
"""Overriding str(*) method to make more human-readable."""
return stringify_object(self)
96 changes: 96 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2206,3 +2206,99 @@ def __init__(
self.skip_model_validation = skip_model_validation
self.source_uri = source_uri
self.config_name = config_name


class BaseDeploymentConfigDataHolder(JumpStartDataHolderType):
"""Base class for Deployment Config Data."""

def _convert_to_pascal_case(self, attr_name: str) -> str:
"""Converts a snake_case attribute name into a camelCased string.

Args:
attr_name (str): The snake_case attribute name.
Returns:
str: The PascalCased attribute name.
"""
return attr_name.replace("_", " ").title().replace(" ", "")

def to_json(self) -> Dict[str, Any]:
"""Represents ``This`` object as JSON."""
json_obj = {}
for att in self.__slots__:
if hasattr(self, att):
cur_val = getattr(self, att)
att = self._convert_to_pascal_case(att)
if issubclass(type(cur_val), JumpStartDataHolderType):
json_obj[att] = cur_val.to_json()
elif isinstance(cur_val, list):
json_obj[att] = []
for obj in cur_val:
if issubclass(type(obj), JumpStartDataHolderType):

@evakravievakraviApr 23, 2024

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

this logic's really complicated. if you can find a way to reduce indentation level, that'd improve readability

json_obj[att].append(obj.to_json())
else:
json_obj[att].append(obj)
elif isinstance(cur_val, dict):
json_obj[att] = {}
for key, val in cur_val.items():
if issubclass(type(val), JumpStartDataHolderType):
json_obj[att][self._convert_to_pascal_case(key)] = val.to_json()
else:
json_obj[att][key] = val
else:
json_obj[att] = cur_val
return json_obj


class DeploymentConfig(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config."""

__slots__ = [
"model_data_download_timeout",
"container_startup_health_check_timeout",
"image_uri",
"model_data",
"instance_type",
"environment",
"compute_resource_requirements",
]

def __init__(
self, init_kwargs: JumpStartModelInitKwargs, deploy_kwargs: JumpStartModelDeployKwargs
):
"""Instantiates DeploymentConfig object."""
if init_kwargs is not None:
self.image_uri = init_kwargs.image_uri
self.model_data = init_kwargs.model_data
self.instance_type = init_kwargs.instance_type
self.environment = init_kwargs.env
if init_kwargs.resources is not None:
self.compute_resource_requirements = (
init_kwargs.resources.get_compute_resource_requirements()
)
if deploy_kwargs is not None:
self.model_data_download_timeout = deploy_kwargs.model_data_download_timeout
self.container_startup_health_check_timeout = (
deploy_kwargs.container_startup_health_check_timeout
)


class DeploymentConfigMetadata(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config Metadata"""

__slots__ = [
"config_name",
"benchmark_metrics",
"deployment_config",
]

def __init__(
self,
config_name: str,
benchmark_metrics: List[JumpStartBenchmarkStat],
init_kwargs: JumpStartModelInitKwargs,
deploy_kwargs: JumpStartModelDeployKwargs,
):
"""Instantiates DeploymentConfigMetadata object."""
self.config_name = config_name
self.benchmark_metrics = benchmark_metrics
self.deployment_config = DeploymentConfig(init_kwargs, deploy_kwargs)
41 changes: 41 additions & 0 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -999,3 +999,44 @@ def get_jumpstart_configs(
if metadata_configs
else {}
)


def extract_metrics_from_deployment_configs(
deployment_configs: List[Dict[str, Any]], config_name: str
) -> Dict[str, List[str]]:
"""Extracts metrics from deployment configs.

Args:
deployment_configs (list[dict[str, Any]]): List of deployment configs.
config_name (str): The name of the deployment config use by the model.
"""

data = {"Config Name": [], "Instance Type": [], "Selected": []}

for index, deployment_config in enumerate(deployment_configs):
if deployment_config.get("DeploymentConfig") is None:
continue

benchmark_metrics = deployment_config.get("BenchmarkMetrics")
if benchmark_metrics is not None:
data["Config Name"].append(deployment_config.get("ConfigName"))
data["Instance Type"].append(
deployment_config.get("DeploymentConfig").get("InstanceType")
Comment thread
makungaj1 marked this conversation as resolved.
)
data["Selected"].append(
"Yes"
if (config_name is not None and config_name == deployment_config.get("ConfigName"))
else "No"
)

if index == 0:
for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
data[column_name] = []

for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
if column_name in data.keys():
data[column_name].append(benchmark_metric.get("value"))

return data
14 changes: 13 additions & 1 deletion src/sagemaker/serve/builder/jumpstart_builder.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,7 +16,7 @@
import copy
from abc import ABC, abstractmethod
from datetime import datetime, timedelta
from typing import Type
from typing import Type, Any, List, Dict
import logging

from sagemaker.model import Model
Expand DownExpand Up@@ -431,6 +431,18 @@ def tune_for_tgi_jumpstart(self, max_tuning_duration: int = 1800):
sharded_supported=sharded_supported, max_tuning_duration=max_tuning_duration
)

def display_benchmark_metrics(self):
"""Display Markdown Benchmark Metrics for deployment configs."""
self.pysdk_model.display_benchmark_metrics()

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model in the current region.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self.pysdk_model.list_deployment_configs()

def _build_for_jumpstart(self):
"""Placeholder docstring"""
# we do not pickle for jumpstart. set to none
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
109 changes: 106 additions & 3 deletions src/sagemaker/jumpstart/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,7 +14,9 @@

from __future__ import absolute_import

from typing import Dict, List, Optional, Union
from functools import lru_cache
from typing import Dict, List, Optional, Union, Any
import pandas as pd
from botocore.exceptions import ClientError

from sagemaker import payloads
Expand All@@ -36,14 +38,21 @@
get_init_kwargs,
get_register_kwargs,
)
from sagemaker.jumpstart.types import JumpStartSerializablePayload
from sagemaker.jumpstart.types import (
JumpStartSerializablePayload,
DeploymentConfigMetadata,
JumpStartBenchmarkStat,
JumpStartMetadataConfig,
)
from sagemaker.jumpstart.utils import (
validate_model_id_and_get_type,
verify_model_region_and_return_specs,
get_jumpstart_configs,
extract_metrics_from_deployment_configs,
)
from sagemaker.jumpstart.constants import JUMPSTART_LOGGER
from sagemaker.jumpstart.enums import JumpStartModelType
from sagemaker.utils import stringify_object, format_tags, Tags
from sagemaker.utils import stringify_object, format_tags, Tags, get_instance_rate_per_hour
from sagemaker.model import (
Model,
ModelPackage,
Expand DownExpand Up@@ -352,6 +361,18 @@ def _validate_model_id_and_type():
self.model_package_arn = model_init_kwargs.model_package_arn
self.init_kwargs = model_init_kwargs.to_kwargs_dict(False)

metadata_configs = get_jumpstart_configs(
region=self.region,
model_id=self.model_id,
model_version=self.model_version,
sagemaker_session=self.sagemaker_session,
model_type=self.model_type,
)
self._deployment_configs = [
self._convert_to_deployment_config_metadata(config_name, config)
for config_name, config in metadata_configs.items()
]

def log_subscription_warning(self) -> None:
"""Log message prompting the customer to subscribe to the proprietary model."""
subscription_link = verify_model_region_and_return_specs(
Expand DownExpand Up@@ -420,6 +441,27 @@ def set_deployment_config(self, config_name: Optional[str]) -> None:
model_id=self.model_id, model_version=self.model_version, config_name=config_name
)

@property
def benchmark_metrics(self) -> pd.DataFrame:
"""Benchmark Metrics for deployment configs

Returns:
Metrics: Pandas DataFrame object.
"""
return pd.DataFrame(self._get_benchmark_data(self.config_name))

def display_benchmark_metrics(self) -> None:
"""Display Benchmark Metrics for deployment configs."""
print(self.benchmark_metrics.to_markdown())

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self._deployment_configs

def _create_sagemaker_model(
self,
instance_type=None,
Expand DownExpand Up@@ -808,6 +850,67 @@ def register_deploy_wrapper(*args, **kwargs):

return model_package

@lru_cache
def _get_benchmark_data(self, config_name: str) -> Dict[str, List[str]]:
"""Constructs deployment configs benchmark data.

Args:
config_name (str): The name of the selected deployment config.
Returns:
Dict[str, List[str]]: Deployment config benchmark data.
"""
return extract_metrics_from_deployment_configs(
self._deployment_configs,
config_name,
)

def _convert_to_deployment_config_metadata(
self, config_name: str, metadata_config: JumpStartMetadataConfig
) -> Dict[str, Any]:
"""Retrieve deployment config for config name.

Args:
config_name (str): Name of deployment config.
metadata_config (JumpStartMetadataConfig): Metadata config for deployment config.
Returns:
A deployment metadata config for config name (dict[str, Any]).
"""
default_inference_instance_type = metadata_config.resolved_config.get(
"default_inference_instance_type"
)

instance_rate = get_instance_rate_per_hour(
instance_type=default_inference_instance_type, region=self.region
)

benchmark_metrics = (
metadata_config.benchmark_metrics.get(default_inference_instance_type)
if metadata_config.benchmark_metrics is not None
else None
)
if instance_rate is not None:
if benchmark_metrics is not None:
benchmark_metrics.append(JumpStartBenchmarkStat(instance_rate))
else:
benchmark_metrics = [JumpStartBenchmarkStat(instance_rate)]

init_kwargs = get_init_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)
deploy_kwargs = get_deploy_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)

deployment_config_metadata = DeploymentConfigMetadata(
config_name, benchmark_metrics, init_kwargs, deploy_kwargs
)

return deployment_config_metadata.to_json()

def __str__(self) -> str:
"""Overriding str(*) method to make more human-readable."""
return stringify_object(self)
96 changes: 96 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2206,3 +2206,99 @@ def __init__(
self.skip_model_validation = skip_model_validation
self.source_uri = source_uri
self.config_name = config_name


class BaseDeploymentConfigDataHolder(JumpStartDataHolderType):
"""Base class for Deployment Config Data."""

def _convert_to_pascal_case(self, attr_name: str) -> str:
"""Converts a snake_case attribute name into a camelCased string.

Args:
attr_name (str): The snake_case attribute name.
Returns:
str: The PascalCased attribute name.
"""
return attr_name.replace("_", " ").title().replace(" ", "")

def to_json(self) -> Dict[str, Any]:
"""Represents ``This`` object as JSON."""
json_obj = {}
for att in self.__slots__:
if hasattr(self, att):
cur_val = getattr(self, att)
att = self._convert_to_pascal_case(att)
if issubclass(type(cur_val), JumpStartDataHolderType):
json_obj[att] = cur_val.to_json()
elif isinstance(cur_val, list):
json_obj[att] = []
for obj in cur_val:
if issubclass(type(obj), JumpStartDataHolderType):

@evakravievakraviApr 23, 2024

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

this logic's really complicated. if you can find a way to reduce indentation level, that'd improve readability

json_obj[att].append(obj.to_json())
else:
json_obj[att].append(obj)
elif isinstance(cur_val, dict):
json_obj[att] = {}
for key, val in cur_val.items():
if issubclass(type(val), JumpStartDataHolderType):
json_obj[att][self._convert_to_pascal_case(key)] = val.to_json()
else:
json_obj[att][key] = val
else:
json_obj[att] = cur_val
return json_obj


class DeploymentConfig(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config."""

__slots__ = [
"model_data_download_timeout",
"container_startup_health_check_timeout",
"image_uri",
"model_data",
"instance_type",
"environment",
"compute_resource_requirements",
]

def __init__(
self, init_kwargs: JumpStartModelInitKwargs, deploy_kwargs: JumpStartModelDeployKwargs
):
"""Instantiates DeploymentConfig object."""
if init_kwargs is not None:
self.image_uri = init_kwargs.image_uri
self.model_data = init_kwargs.model_data
self.instance_type = init_kwargs.instance_type
self.environment = init_kwargs.env
if init_kwargs.resources is not None:
self.compute_resource_requirements = (
init_kwargs.resources.get_compute_resource_requirements()
)
if deploy_kwargs is not None:
self.model_data_download_timeout = deploy_kwargs.model_data_download_timeout
self.container_startup_health_check_timeout = (
deploy_kwargs.container_startup_health_check_timeout
)


class DeploymentConfigMetadata(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config Metadata"""

__slots__ = [
"config_name",
"benchmark_metrics",
"deployment_config",
]

def __init__(
self,
config_name: str,
benchmark_metrics: List[JumpStartBenchmarkStat],
init_kwargs: JumpStartModelInitKwargs,
deploy_kwargs: JumpStartModelDeployKwargs,
):
"""Instantiates DeploymentConfigMetadata object."""
self.config_name = config_name
self.benchmark_metrics = benchmark_metrics
self.deployment_config = DeploymentConfig(init_kwargs, deploy_kwargs)
41 changes: 41 additions & 0 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -999,3 +999,44 @@ def get_jumpstart_configs(
if metadata_configs
else {}
)


def extract_metrics_from_deployment_configs(
deployment_configs: List[Dict[str, Any]], config_name: str
) -> Dict[str, List[str]]:
"""Extracts metrics from deployment configs.

Args:
deployment_configs (list[dict[str, Any]]): List of deployment configs.
config_name (str): The name of the deployment config use by the model.
"""

data = {"Config Name": [], "Instance Type": [], "Selected": []}

for index, deployment_config in enumerate(deployment_configs):
if deployment_config.get("DeploymentConfig") is None:
continue

benchmark_metrics = deployment_config.get("BenchmarkMetrics")
if benchmark_metrics is not None:
data["Config Name"].append(deployment_config.get("ConfigName"))
data["Instance Type"].append(
deployment_config.get("DeploymentConfig").get("InstanceType")
Comment thread
makungaj1 marked this conversation as resolved.
)
data["Selected"].append(
"Yes"
if (config_name is not None and config_name == deployment_config.get("ConfigName"))
else "No"
)

if index == 0:
for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
data[column_name] = []

for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
if column_name in data.keys():
data[column_name].append(benchmark_metric.get("value"))

return data
14 changes: 13 additions & 1 deletion src/sagemaker/serve/builder/jumpstart_builder.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,7 +16,7 @@
import copy
from abc import ABC, abstractmethod
from datetime import datetime, timedelta
from typing import Type
from typing import Type, Any, List, Dict
import logging

from sagemaker.model import Model
Expand DownExpand Up@@ -431,6 +431,18 @@ def tune_for_tgi_jumpstart(self, max_tuning_duration: int = 1800):
sharded_supported=sharded_supported, max_tuning_duration=max_tuning_duration
)

def display_benchmark_metrics(self):
"""Display Markdown Benchmark Metrics for deployment configs."""
self.pysdk_model.display_benchmark_metrics()

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model in the current region.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self.pysdk_model.list_deployment_configs()

def _build_for_jumpstart(self):
"""Placeholder docstring"""
# we do not pickle for jumpstart. set to none
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
109 changes: 106 additions & 3 deletions src/sagemaker/jumpstart/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,7 +14,9 @@

from __future__ import absolute_import

from typing import Dict, List, Optional, Union
from functools import lru_cache
from typing import Dict, List, Optional, Union, Any
import pandas as pd
from botocore.exceptions import ClientError

from sagemaker import payloads
Expand All@@ -36,14 +38,21 @@
get_init_kwargs,
get_register_kwargs,
)
from sagemaker.jumpstart.types import JumpStartSerializablePayload
from sagemaker.jumpstart.types import (
JumpStartSerializablePayload,
DeploymentConfigMetadata,
JumpStartBenchmarkStat,
JumpStartMetadataConfig,
)
from sagemaker.jumpstart.utils import (
validate_model_id_and_get_type,
verify_model_region_and_return_specs,
get_jumpstart_configs,
extract_metrics_from_deployment_configs,
)
from sagemaker.jumpstart.constants import JUMPSTART_LOGGER
from sagemaker.jumpstart.enums import JumpStartModelType
from sagemaker.utils import stringify_object, format_tags, Tags
from sagemaker.utils import stringify_object, format_tags, Tags, get_instance_rate_per_hour
from sagemaker.model import (
Model,
ModelPackage,
Expand DownExpand Up@@ -352,6 +361,18 @@ def _validate_model_id_and_type():
self.model_package_arn = model_init_kwargs.model_package_arn
self.init_kwargs = model_init_kwargs.to_kwargs_dict(False)

metadata_configs = get_jumpstart_configs(
region=self.region,
model_id=self.model_id,
model_version=self.model_version,
sagemaker_session=self.sagemaker_session,
model_type=self.model_type,
)
self._deployment_configs = [
self._convert_to_deployment_config_metadata(config_name, config)
for config_name, config in metadata_configs.items()
]

def log_subscription_warning(self) -> None:
"""Log message prompting the customer to subscribe to the proprietary model."""
subscription_link = verify_model_region_and_return_specs(
Expand DownExpand Up@@ -420,6 +441,27 @@ def set_deployment_config(self, config_name: Optional[str]) -> None:
model_id=self.model_id, model_version=self.model_version, config_name=config_name
)

@property
def benchmark_metrics(self) -> pd.DataFrame:
"""Benchmark Metrics for deployment configs

Returns:
Metrics: Pandas DataFrame object.
"""
return pd.DataFrame(self._get_benchmark_data(self.config_name))

def display_benchmark_metrics(self) -> None:
"""Display Benchmark Metrics for deployment configs."""
print(self.benchmark_metrics.to_markdown())

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self._deployment_configs

def _create_sagemaker_model(
self,
instance_type=None,
Expand DownExpand Up@@ -808,6 +850,67 @@ def register_deploy_wrapper(*args, **kwargs):

return model_package

@lru_cache
def _get_benchmark_data(self, config_name: str) -> Dict[str, List[str]]:
"""Constructs deployment configs benchmark data.

Args:
config_name (str): The name of the selected deployment config.
Returns:
Dict[str, List[str]]: Deployment config benchmark data.
"""
return extract_metrics_from_deployment_configs(
self._deployment_configs,
config_name,
)

def _convert_to_deployment_config_metadata(
self, config_name: str, metadata_config: JumpStartMetadataConfig
) -> Dict[str, Any]:
"""Retrieve deployment config for config name.

Args:
config_name (str): Name of deployment config.
metadata_config (JumpStartMetadataConfig): Metadata config for deployment config.
Returns:
A deployment metadata config for config name (dict[str, Any]).
"""
default_inference_instance_type = metadata_config.resolved_config.get(
"default_inference_instance_type"
)

instance_rate = get_instance_rate_per_hour(
instance_type=default_inference_instance_type, region=self.region
)

benchmark_metrics = (
metadata_config.benchmark_metrics.get(default_inference_instance_type)
if metadata_config.benchmark_metrics is not None
else None
)
if instance_rate is not None:
if benchmark_metrics is not None:
benchmark_metrics.append(JumpStartBenchmarkStat(instance_rate))
else:
benchmark_metrics = [JumpStartBenchmarkStat(instance_rate)]

init_kwargs = get_init_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)
deploy_kwargs = get_deploy_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)

deployment_config_metadata = DeploymentConfigMetadata(
config_name, benchmark_metrics, init_kwargs, deploy_kwargs
)

return deployment_config_metadata.to_json()

def __str__(self) -> str:
"""Overriding str(*) method to make more human-readable."""
return stringify_object(self)
96 changes: 96 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2206,3 +2206,99 @@ def __init__(
self.skip_model_validation = skip_model_validation
self.source_uri = source_uri
self.config_name = config_name


class BaseDeploymentConfigDataHolder(JumpStartDataHolderType):
"""Base class for Deployment Config Data."""

def _convert_to_pascal_case(self, attr_name: str) -> str:
"""Converts a snake_case attribute name into a camelCased string.

Args:
attr_name (str): The snake_case attribute name.
Returns:
str: The PascalCased attribute name.
"""
return attr_name.replace("_", " ").title().replace(" ", "")

def to_json(self) -> Dict[str, Any]:
"""Represents ``This`` object as JSON."""
json_obj = {}
for att in self.__slots__:
if hasattr(self, att):
cur_val = getattr(self, att)
att = self._convert_to_pascal_case(att)
if issubclass(type(cur_val), JumpStartDataHolderType):
json_obj[att] = cur_val.to_json()
elif isinstance(cur_val, list):
json_obj[att] = []
for obj in cur_val:
if issubclass(type(obj), JumpStartDataHolderType):

@evakravievakraviApr 23, 2024

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

this logic's really complicated. if you can find a way to reduce indentation level, that'd improve readability

json_obj[att].append(obj.to_json())
else:
json_obj[att].append(obj)
elif isinstance(cur_val, dict):
json_obj[att] = {}
for key, val in cur_val.items():
if issubclass(type(val), JumpStartDataHolderType):
json_obj[att][self._convert_to_pascal_case(key)] = val.to_json()
else:
json_obj[att][key] = val
else:
json_obj[att] = cur_val
return json_obj


class DeploymentConfig(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config."""

__slots__ = [
"model_data_download_timeout",
"container_startup_health_check_timeout",
"image_uri",
"model_data",
"instance_type",
"environment",
"compute_resource_requirements",
]

def __init__(
self, init_kwargs: JumpStartModelInitKwargs, deploy_kwargs: JumpStartModelDeployKwargs
):
"""Instantiates DeploymentConfig object."""
if init_kwargs is not None:
self.image_uri = init_kwargs.image_uri
self.model_data = init_kwargs.model_data
self.instance_type = init_kwargs.instance_type
self.environment = init_kwargs.env
if init_kwargs.resources is not None:
self.compute_resource_requirements = (
init_kwargs.resources.get_compute_resource_requirements()
)
if deploy_kwargs is not None:
self.model_data_download_timeout = deploy_kwargs.model_data_download_timeout
self.container_startup_health_check_timeout = (
deploy_kwargs.container_startup_health_check_timeout
)


class DeploymentConfigMetadata(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config Metadata"""

__slots__ = [
"config_name",
"benchmark_metrics",
"deployment_config",
]

def __init__(
self,
config_name: str,
benchmark_metrics: List[JumpStartBenchmarkStat],
init_kwargs: JumpStartModelInitKwargs,
deploy_kwargs: JumpStartModelDeployKwargs,
):
"""Instantiates DeploymentConfigMetadata object."""
self.config_name = config_name
self.benchmark_metrics = benchmark_metrics
self.deployment_config = DeploymentConfig(init_kwargs, deploy_kwargs)
41 changes: 41 additions & 0 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -999,3 +999,44 @@ def get_jumpstart_configs(
if metadata_configs
else {}
)


def extract_metrics_from_deployment_configs(
deployment_configs: List[Dict[str, Any]], config_name: str
) -> Dict[str, List[str]]:
"""Extracts metrics from deployment configs.

Args:
deployment_configs (list[dict[str, Any]]): List of deployment configs.
config_name (str): The name of the deployment config use by the model.
"""

data = {"Config Name": [], "Instance Type": [], "Selected": []}

for index, deployment_config in enumerate(deployment_configs):
if deployment_config.get("DeploymentConfig") is None:
continue

benchmark_metrics = deployment_config.get("BenchmarkMetrics")
if benchmark_metrics is not None:
data["Config Name"].append(deployment_config.get("ConfigName"))
data["Instance Type"].append(
deployment_config.get("DeploymentConfig").get("InstanceType")
Comment thread
makungaj1 marked this conversation as resolved.
)
data["Selected"].append(
"Yes"
if (config_name is not None and config_name == deployment_config.get("ConfigName"))
else "No"
)

if index == 0:
for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
data[column_name] = []

for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
if column_name in data.keys():
data[column_name].append(benchmark_metric.get("value"))

return data
14 changes: 13 additions & 1 deletion src/sagemaker/serve/builder/jumpstart_builder.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,7 +16,7 @@
import copy
from abc import ABC, abstractmethod
from datetime import datetime, timedelta
from typing import Type
from typing import Type, Any, List, Dict
import logging

from sagemaker.model import Model
Expand DownExpand Up@@ -431,6 +431,18 @@ def tune_for_tgi_jumpstart(self, max_tuning_duration: int = 1800):
sharded_supported=sharded_supported, max_tuning_duration=max_tuning_duration
)

def display_benchmark_metrics(self):
"""Display Markdown Benchmark Metrics for deployment configs."""
self.pysdk_model.display_benchmark_metrics()

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model in the current region.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self.pysdk_model.list_deployment_configs()

def _build_for_jumpstart(self):
"""Placeholder docstring"""
# we do not pickle for jumpstart. set to none
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
109 changes: 106 additions & 3 deletions src/sagemaker/jumpstart/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,7 +14,9 @@

from __future__ import absolute_import

from typing import Dict, List, Optional, Union
from functools import lru_cache
from typing import Dict, List, Optional, Union, Any
import pandas as pd
from botocore.exceptions import ClientError

from sagemaker import payloads
Expand All@@ -36,14 +38,21 @@
get_init_kwargs,
get_register_kwargs,
)
from sagemaker.jumpstart.types import JumpStartSerializablePayload
from sagemaker.jumpstart.types import (
JumpStartSerializablePayload,
DeploymentConfigMetadata,
JumpStartBenchmarkStat,
JumpStartMetadataConfig,
)
from sagemaker.jumpstart.utils import (
validate_model_id_and_get_type,
verify_model_region_and_return_specs,
get_jumpstart_configs,
extract_metrics_from_deployment_configs,
)
from sagemaker.jumpstart.constants import JUMPSTART_LOGGER
from sagemaker.jumpstart.enums import JumpStartModelType
from sagemaker.utils import stringify_object, format_tags, Tags
from sagemaker.utils import stringify_object, format_tags, Tags, get_instance_rate_per_hour
from sagemaker.model import (
Model,
ModelPackage,
Expand DownExpand Up@@ -352,6 +361,18 @@ def _validate_model_id_and_type():
self.model_package_arn = model_init_kwargs.model_package_arn
self.init_kwargs = model_init_kwargs.to_kwargs_dict(False)

metadata_configs = get_jumpstart_configs(
region=self.region,
model_id=self.model_id,
model_version=self.model_version,
sagemaker_session=self.sagemaker_session,
model_type=self.model_type,
)
self._deployment_configs = [
self._convert_to_deployment_config_metadata(config_name, config)
for config_name, config in metadata_configs.items()
]

def log_subscription_warning(self) -> None:
"""Log message prompting the customer to subscribe to the proprietary model."""
subscription_link = verify_model_region_and_return_specs(
Expand DownExpand Up@@ -420,6 +441,27 @@ def set_deployment_config(self, config_name: Optional[str]) -> None:
model_id=self.model_id, model_version=self.model_version, config_name=config_name
)

@property
def benchmark_metrics(self) -> pd.DataFrame:
"""Benchmark Metrics for deployment configs

Returns:
Metrics: Pandas DataFrame object.
"""
return pd.DataFrame(self._get_benchmark_data(self.config_name))

def display_benchmark_metrics(self) -> None:
"""Display Benchmark Metrics for deployment configs."""
print(self.benchmark_metrics.to_markdown())

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self._deployment_configs

def _create_sagemaker_model(
self,
instance_type=None,
Expand DownExpand Up@@ -808,6 +850,67 @@ def register_deploy_wrapper(*args, **kwargs):

return model_package

@lru_cache
def _get_benchmark_data(self, config_name: str) -> Dict[str, List[str]]:
"""Constructs deployment configs benchmark data.

Args:
config_name (str): The name of the selected deployment config.
Returns:
Dict[str, List[str]]: Deployment config benchmark data.
"""
return extract_metrics_from_deployment_configs(
self._deployment_configs,
config_name,
)

def _convert_to_deployment_config_metadata(
self, config_name: str, metadata_config: JumpStartMetadataConfig
) -> Dict[str, Any]:
"""Retrieve deployment config for config name.

Args:
config_name (str): Name of deployment config.
metadata_config (JumpStartMetadataConfig): Metadata config for deployment config.
Returns:
A deployment metadata config for config name (dict[str, Any]).
"""
default_inference_instance_type = metadata_config.resolved_config.get(
"default_inference_instance_type"
)

instance_rate = get_instance_rate_per_hour(
instance_type=default_inference_instance_type, region=self.region
)

benchmark_metrics = (
metadata_config.benchmark_metrics.get(default_inference_instance_type)
if metadata_config.benchmark_metrics is not None
else None
)
if instance_rate is not None:
if benchmark_metrics is not None:
benchmark_metrics.append(JumpStartBenchmarkStat(instance_rate))
else:
benchmark_metrics = [JumpStartBenchmarkStat(instance_rate)]

init_kwargs = get_init_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)
deploy_kwargs = get_deploy_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)

deployment_config_metadata = DeploymentConfigMetadata(
config_name, benchmark_metrics, init_kwargs, deploy_kwargs
)

return deployment_config_metadata.to_json()

def __str__(self) -> str:
"""Overriding str(*) method to make more human-readable."""
return stringify_object(self)
96 changes: 96 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2206,3 +2206,99 @@ def __init__(
self.skip_model_validation = skip_model_validation
self.source_uri = source_uri
self.config_name = config_name


class BaseDeploymentConfigDataHolder(JumpStartDataHolderType):
"""Base class for Deployment Config Data."""

def _convert_to_pascal_case(self, attr_name: str) -> str:
"""Converts a snake_case attribute name into a camelCased string.

Args:
attr_name (str): The snake_case attribute name.
Returns:
str: The PascalCased attribute name.
"""
return attr_name.replace("_", " ").title().replace(" ", "")

def to_json(self) -> Dict[str, Any]:
"""Represents ``This`` object as JSON."""
json_obj = {}
for att in self.__slots__:
if hasattr(self, att):
cur_val = getattr(self, att)
att = self._convert_to_pascal_case(att)
if issubclass(type(cur_val), JumpStartDataHolderType):
json_obj[att] = cur_val.to_json()
elif isinstance(cur_val, list):
json_obj[att] = []
for obj in cur_val:
if issubclass(type(obj), JumpStartDataHolderType):

@evakravievakraviApr 23, 2024

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

this logic's really complicated. if you can find a way to reduce indentation level, that'd improve readability

json_obj[att].append(obj.to_json())
else:
json_obj[att].append(obj)
elif isinstance(cur_val, dict):
json_obj[att] = {}
for key, val in cur_val.items():
if issubclass(type(val), JumpStartDataHolderType):
json_obj[att][self._convert_to_pascal_case(key)] = val.to_json()
else:
json_obj[att][key] = val
else:
json_obj[att] = cur_val
return json_obj


class DeploymentConfig(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config."""

__slots__ = [
"model_data_download_timeout",
"container_startup_health_check_timeout",
"image_uri",
"model_data",
"instance_type",
"environment",
"compute_resource_requirements",
]

def __init__(
self, init_kwargs: JumpStartModelInitKwargs, deploy_kwargs: JumpStartModelDeployKwargs
):
"""Instantiates DeploymentConfig object."""
if init_kwargs is not None:
self.image_uri = init_kwargs.image_uri
self.model_data = init_kwargs.model_data
self.instance_type = init_kwargs.instance_type
self.environment = init_kwargs.env
if init_kwargs.resources is not None:
self.compute_resource_requirements = (
init_kwargs.resources.get_compute_resource_requirements()
)
if deploy_kwargs is not None:
self.model_data_download_timeout = deploy_kwargs.model_data_download_timeout
self.container_startup_health_check_timeout = (
deploy_kwargs.container_startup_health_check_timeout
)


class DeploymentConfigMetadata(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config Metadata"""

__slots__ = [
"config_name",
"benchmark_metrics",
"deployment_config",
]

def __init__(
self,
config_name: str,
benchmark_metrics: List[JumpStartBenchmarkStat],
init_kwargs: JumpStartModelInitKwargs,
deploy_kwargs: JumpStartModelDeployKwargs,
):
"""Instantiates DeploymentConfigMetadata object."""
self.config_name = config_name
self.benchmark_metrics = benchmark_metrics
self.deployment_config = DeploymentConfig(init_kwargs, deploy_kwargs)
41 changes: 41 additions & 0 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -999,3 +999,44 @@ def get_jumpstart_configs(
if metadata_configs
else {}
)


def extract_metrics_from_deployment_configs(
deployment_configs: List[Dict[str, Any]], config_name: str
) -> Dict[str, List[str]]:
"""Extracts metrics from deployment configs.

Args:
deployment_configs (list[dict[str, Any]]): List of deployment configs.
config_name (str): The name of the deployment config use by the model.
"""

data = {"Config Name": [], "Instance Type": [], "Selected": []}

for index, deployment_config in enumerate(deployment_configs):
if deployment_config.get("DeploymentConfig") is None:
continue

benchmark_metrics = deployment_config.get("BenchmarkMetrics")
if benchmark_metrics is not None:
data["Config Name"].append(deployment_config.get("ConfigName"))
data["Instance Type"].append(
deployment_config.get("DeploymentConfig").get("InstanceType")
Comment thread
makungaj1 marked this conversation as resolved.
)
data["Selected"].append(
"Yes"
if (config_name is not None and config_name == deployment_config.get("ConfigName"))
else "No"
)

if index == 0:
for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
data[column_name] = []

for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
if column_name in data.keys():
data[column_name].append(benchmark_metric.get("value"))

return data
14 changes: 13 additions & 1 deletion src/sagemaker/serve/builder/jumpstart_builder.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,7 +16,7 @@
import copy
from abc import ABC, abstractmethod
from datetime import datetime, timedelta
from typing import Type
from typing import Type, Any, List, Dict
import logging

from sagemaker.model import Model
Expand DownExpand Up@@ -431,6 +431,18 @@ def tune_for_tgi_jumpstart(self, max_tuning_duration: int = 1800):
sharded_supported=sharded_supported, max_tuning_duration=max_tuning_duration
)

def display_benchmark_metrics(self):
"""Display Markdown Benchmark Metrics for deployment configs."""
self.pysdk_model.display_benchmark_metrics()

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model in the current region.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self.pysdk_model.list_deployment_configs()

def _build_for_jumpstart(self):
"""Placeholder docstring"""
# we do not pickle for jumpstart. set to none
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
109 changes: 106 additions & 3 deletions src/sagemaker/jumpstart/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,7 +14,9 @@

from __future__ import absolute_import

from typing import Dict, List, Optional, Union
from functools import lru_cache
from typing import Dict, List, Optional, Union, Any
import pandas as pd
from botocore.exceptions import ClientError

from sagemaker import payloads
Expand All@@ -36,14 +38,21 @@
get_init_kwargs,
get_register_kwargs,
)
from sagemaker.jumpstart.types import JumpStartSerializablePayload
from sagemaker.jumpstart.types import (
JumpStartSerializablePayload,
DeploymentConfigMetadata,
JumpStartBenchmarkStat,
JumpStartMetadataConfig,
)
from sagemaker.jumpstart.utils import (
validate_model_id_and_get_type,
verify_model_region_and_return_specs,
get_jumpstart_configs,
extract_metrics_from_deployment_configs,
)
from sagemaker.jumpstart.constants import JUMPSTART_LOGGER
from sagemaker.jumpstart.enums import JumpStartModelType
from sagemaker.utils import stringify_object, format_tags, Tags
from sagemaker.utils import stringify_object, format_tags, Tags, get_instance_rate_per_hour
from sagemaker.model import (
Model,
ModelPackage,
Expand DownExpand Up@@ -352,6 +361,18 @@ def _validate_model_id_and_type():
self.model_package_arn = model_init_kwargs.model_package_arn
self.init_kwargs = model_init_kwargs.to_kwargs_dict(False)

metadata_configs = get_jumpstart_configs(
region=self.region,
model_id=self.model_id,
model_version=self.model_version,
sagemaker_session=self.sagemaker_session,
model_type=self.model_type,
)
self._deployment_configs = [
self._convert_to_deployment_config_metadata(config_name, config)
for config_name, config in metadata_configs.items()
]

def log_subscription_warning(self) -> None:
"""Log message prompting the customer to subscribe to the proprietary model."""
subscription_link = verify_model_region_and_return_specs(
Expand DownExpand Up@@ -420,6 +441,27 @@ def set_deployment_config(self, config_name: Optional[str]) -> None:
model_id=self.model_id, model_version=self.model_version, config_name=config_name
)

@property
def benchmark_metrics(self) -> pd.DataFrame:
"""Benchmark Metrics for deployment configs

Returns:
Metrics: Pandas DataFrame object.
"""
return pd.DataFrame(self._get_benchmark_data(self.config_name))

def display_benchmark_metrics(self) -> None:
"""Display Benchmark Metrics for deployment configs."""
print(self.benchmark_metrics.to_markdown())

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self._deployment_configs

def _create_sagemaker_model(
self,
instance_type=None,
Expand DownExpand Up@@ -808,6 +850,67 @@ def register_deploy_wrapper(*args, **kwargs):

return model_package

@lru_cache
def _get_benchmark_data(self, config_name: str) -> Dict[str, List[str]]:
"""Constructs deployment configs benchmark data.

Args:
config_name (str): The name of the selected deployment config.
Returns:
Dict[str, List[str]]: Deployment config benchmark data.
"""
return extract_metrics_from_deployment_configs(
self._deployment_configs,
config_name,
)

def _convert_to_deployment_config_metadata(
self, config_name: str, metadata_config: JumpStartMetadataConfig
) -> Dict[str, Any]:
"""Retrieve deployment config for config name.

Args:
config_name (str): Name of deployment config.
metadata_config (JumpStartMetadataConfig): Metadata config for deployment config.
Returns:
A deployment metadata config for config name (dict[str, Any]).
"""
default_inference_instance_type = metadata_config.resolved_config.get(
"default_inference_instance_type"
)

instance_rate = get_instance_rate_per_hour(
instance_type=default_inference_instance_type, region=self.region
)

benchmark_metrics = (
metadata_config.benchmark_metrics.get(default_inference_instance_type)
if metadata_config.benchmark_metrics is not None
else None
)
if instance_rate is not None:
if benchmark_metrics is not None:
benchmark_metrics.append(JumpStartBenchmarkStat(instance_rate))
else:
benchmark_metrics = [JumpStartBenchmarkStat(instance_rate)]

init_kwargs = get_init_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)
deploy_kwargs = get_deploy_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)

deployment_config_metadata = DeploymentConfigMetadata(
config_name, benchmark_metrics, init_kwargs, deploy_kwargs
)

return deployment_config_metadata.to_json()

def __str__(self) -> str:
"""Overriding str(*) method to make more human-readable."""
return stringify_object(self)
96 changes: 96 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2206,3 +2206,99 @@ def __init__(
self.skip_model_validation = skip_model_validation
self.source_uri = source_uri
self.config_name = config_name


class BaseDeploymentConfigDataHolder(JumpStartDataHolderType):
"""Base class for Deployment Config Data."""

def _convert_to_pascal_case(self, attr_name: str) -> str:
"""Converts a snake_case attribute name into a camelCased string.

Args:
attr_name (str): The snake_case attribute name.
Returns:
str: The PascalCased attribute name.
"""
return attr_name.replace("_", " ").title().replace(" ", "")

def to_json(self) -> Dict[str, Any]:
"""Represents ``This`` object as JSON."""
json_obj = {}
for att in self.__slots__:
if hasattr(self, att):
cur_val = getattr(self, att)
att = self._convert_to_pascal_case(att)
if issubclass(type(cur_val), JumpStartDataHolderType):
json_obj[att] = cur_val.to_json()
elif isinstance(cur_val, list):
json_obj[att] = []
for obj in cur_val:
if issubclass(type(obj), JumpStartDataHolderType):

@evakravievakraviApr 23, 2024

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

this logic's really complicated. if you can find a way to reduce indentation level, that'd improve readability

json_obj[att].append(obj.to_json())
else:
json_obj[att].append(obj)
elif isinstance(cur_val, dict):
json_obj[att] = {}
for key, val in cur_val.items():
if issubclass(type(val), JumpStartDataHolderType):
json_obj[att][self._convert_to_pascal_case(key)] = val.to_json()
else:
json_obj[att][key] = val
else:
json_obj[att] = cur_val
return json_obj


class DeploymentConfig(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config."""

__slots__ = [
"model_data_download_timeout",
"container_startup_health_check_timeout",
"image_uri",
"model_data",
"instance_type",
"environment",
"compute_resource_requirements",
]

def __init__(
self, init_kwargs: JumpStartModelInitKwargs, deploy_kwargs: JumpStartModelDeployKwargs
):
"""Instantiates DeploymentConfig object."""
if init_kwargs is not None:
self.image_uri = init_kwargs.image_uri
self.model_data = init_kwargs.model_data
self.instance_type = init_kwargs.instance_type
self.environment = init_kwargs.env
if init_kwargs.resources is not None:
self.compute_resource_requirements = (
init_kwargs.resources.get_compute_resource_requirements()
)
if deploy_kwargs is not None:
self.model_data_download_timeout = deploy_kwargs.model_data_download_timeout
self.container_startup_health_check_timeout = (
deploy_kwargs.container_startup_health_check_timeout
)


class DeploymentConfigMetadata(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config Metadata"""

__slots__ = [
"config_name",
"benchmark_metrics",
"deployment_config",
]

def __init__(
self,
config_name: str,
benchmark_metrics: List[JumpStartBenchmarkStat],
init_kwargs: JumpStartModelInitKwargs,
deploy_kwargs: JumpStartModelDeployKwargs,
):
"""Instantiates DeploymentConfigMetadata object."""
self.config_name = config_name
self.benchmark_metrics = benchmark_metrics
self.deployment_config = DeploymentConfig(init_kwargs, deploy_kwargs)
41 changes: 41 additions & 0 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -999,3 +999,44 @@ def get_jumpstart_configs(
if metadata_configs
else {}
)


def extract_metrics_from_deployment_configs(
deployment_configs: List[Dict[str, Any]], config_name: str
) -> Dict[str, List[str]]:
"""Extracts metrics from deployment configs.

Args:
deployment_configs (list[dict[str, Any]]): List of deployment configs.
config_name (str): The name of the deployment config use by the model.
"""

data = {"Config Name": [], "Instance Type": [], "Selected": []}

for index, deployment_config in enumerate(deployment_configs):
if deployment_config.get("DeploymentConfig") is None:
continue

benchmark_metrics = deployment_config.get("BenchmarkMetrics")
if benchmark_metrics is not None:
data["Config Name"].append(deployment_config.get("ConfigName"))
data["Instance Type"].append(
deployment_config.get("DeploymentConfig").get("InstanceType")
Comment thread
makungaj1 marked this conversation as resolved.
)
data["Selected"].append(
"Yes"
if (config_name is not None and config_name == deployment_config.get("ConfigName"))
else "No"
)

if index == 0:
for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
data[column_name] = []

for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
if column_name in data.keys():
data[column_name].append(benchmark_metric.get("value"))

return data
14 changes: 13 additions & 1 deletion src/sagemaker/serve/builder/jumpstart_builder.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,7 +16,7 @@
import copy
from abc import ABC, abstractmethod
from datetime import datetime, timedelta
from typing import Type
from typing import Type, Any, List, Dict
import logging

from sagemaker.model import Model
Expand DownExpand Up@@ -431,6 +431,18 @@ def tune_for_tgi_jumpstart(self, max_tuning_duration: int = 1800):
sharded_supported=sharded_supported, max_tuning_duration=max_tuning_duration
)

def display_benchmark_metrics(self):
"""Display Markdown Benchmark Metrics for deployment configs."""
self.pysdk_model.display_benchmark_metrics()

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model in the current region.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self.pysdk_model.list_deployment_configs()

def _build_for_jumpstart(self):
"""Placeholder docstring"""
# we do not pickle for jumpstart. set to none
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
109 changes: 106 additions & 3 deletions src/sagemaker/jumpstart/model.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,7 +14,9 @@

from __future__ import absolute_import

from typing import Dict, List, Optional, Union
from functools import lru_cache
from typing import Dict, List, Optional, Union, Any
import pandas as pd
from botocore.exceptions import ClientError

from sagemaker import payloads
Expand All@@ -36,14 +38,21 @@
get_init_kwargs,
get_register_kwargs,
)
from sagemaker.jumpstart.types import JumpStartSerializablePayload
from sagemaker.jumpstart.types import (
JumpStartSerializablePayload,
DeploymentConfigMetadata,
JumpStartBenchmarkStat,
JumpStartMetadataConfig,
)
from sagemaker.jumpstart.utils import (
validate_model_id_and_get_type,
verify_model_region_and_return_specs,
get_jumpstart_configs,
extract_metrics_from_deployment_configs,
)
from sagemaker.jumpstart.constants import JUMPSTART_LOGGER
from sagemaker.jumpstart.enums import JumpStartModelType
from sagemaker.utils import stringify_object, format_tags, Tags
from sagemaker.utils import stringify_object, format_tags, Tags, get_instance_rate_per_hour
from sagemaker.model import (
Model,
ModelPackage,
Expand DownExpand Up@@ -352,6 +361,18 @@ def _validate_model_id_and_type():
self.model_package_arn = model_init_kwargs.model_package_arn
self.init_kwargs = model_init_kwargs.to_kwargs_dict(False)

metadata_configs = get_jumpstart_configs(
region=self.region,
model_id=self.model_id,
model_version=self.model_version,
sagemaker_session=self.sagemaker_session,
model_type=self.model_type,
)
self._deployment_configs = [
self._convert_to_deployment_config_metadata(config_name, config)
for config_name, config in metadata_configs.items()
]

def log_subscription_warning(self) -> None:
"""Log message prompting the customer to subscribe to the proprietary model."""
subscription_link = verify_model_region_and_return_specs(
Expand DownExpand Up@@ -420,6 +441,27 @@ def set_deployment_config(self, config_name: Optional[str]) -> None:
model_id=self.model_id, model_version=self.model_version, config_name=config_name
)

@property
def benchmark_metrics(self) -> pd.DataFrame:
"""Benchmark Metrics for deployment configs

Returns:
Metrics: Pandas DataFrame object.
"""
return pd.DataFrame(self._get_benchmark_data(self.config_name))

def display_benchmark_metrics(self) -> None:
"""Display Benchmark Metrics for deployment configs."""
print(self.benchmark_metrics.to_markdown())

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self._deployment_configs

def _create_sagemaker_model(
self,
instance_type=None,
Expand DownExpand Up@@ -808,6 +850,67 @@ def register_deploy_wrapper(*args, **kwargs):

return model_package

@lru_cache
def _get_benchmark_data(self, config_name: str) -> Dict[str, List[str]]:
"""Constructs deployment configs benchmark data.

Args:
config_name (str): The name of the selected deployment config.
Returns:
Dict[str, List[str]]: Deployment config benchmark data.
"""
return extract_metrics_from_deployment_configs(
self._deployment_configs,
config_name,
)

def _convert_to_deployment_config_metadata(
self, config_name: str, metadata_config: JumpStartMetadataConfig
) -> Dict[str, Any]:
"""Retrieve deployment config for config name.

Args:
config_name (str): Name of deployment config.
metadata_config (JumpStartMetadataConfig): Metadata config for deployment config.
Returns:
A deployment metadata config for config name (dict[str, Any]).
"""
default_inference_instance_type = metadata_config.resolved_config.get(
"default_inference_instance_type"
)

instance_rate = get_instance_rate_per_hour(
instance_type=default_inference_instance_type, region=self.region
)

benchmark_metrics = (
metadata_config.benchmark_metrics.get(default_inference_instance_type)
if metadata_config.benchmark_metrics is not None
else None
)
if instance_rate is not None:
if benchmark_metrics is not None:
benchmark_metrics.append(JumpStartBenchmarkStat(instance_rate))
else:
benchmark_metrics = [JumpStartBenchmarkStat(instance_rate)]

init_kwargs = get_init_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)
deploy_kwargs = get_deploy_kwargs(
model_id=self.model_id,
instance_type=default_inference_instance_type,
sagemaker_session=self.sagemaker_session,
)

deployment_config_metadata = DeploymentConfigMetadata(
config_name, benchmark_metrics, init_kwargs, deploy_kwargs
)

return deployment_config_metadata.to_json()

def __str__(self) -> str:
"""Overriding str(*) method to make more human-readable."""
return stringify_object(self)
96 changes: 96 additions & 0 deletions src/sagemaker/jumpstart/types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -2206,3 +2206,99 @@ def __init__(
self.skip_model_validation = skip_model_validation
self.source_uri = source_uri
self.config_name = config_name


class BaseDeploymentConfigDataHolder(JumpStartDataHolderType):
"""Base class for Deployment Config Data."""

def _convert_to_pascal_case(self, attr_name: str) -> str:
"""Converts a snake_case attribute name into a camelCased string.

Args:
attr_name (str): The snake_case attribute name.
Returns:
str: The PascalCased attribute name.
"""
return attr_name.replace("_", " ").title().replace(" ", "")

def to_json(self) -> Dict[str, Any]:
"""Represents ``This`` object as JSON."""
json_obj = {}
for att in self.__slots__:
if hasattr(self, att):
cur_val = getattr(self, att)
att = self._convert_to_pascal_case(att)
if issubclass(type(cur_val), JumpStartDataHolderType):
json_obj[att] = cur_val.to_json()
elif isinstance(cur_val, list):
json_obj[att] = []
for obj in cur_val:
if issubclass(type(obj), JumpStartDataHolderType):

@evakravievakraviApr 23, 2024

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

this logic's really complicated. if you can find a way to reduce indentation level, that'd improve readability

json_obj[att].append(obj.to_json())
else:
json_obj[att].append(obj)
elif isinstance(cur_val, dict):
json_obj[att] = {}
for key, val in cur_val.items():
if issubclass(type(val), JumpStartDataHolderType):
json_obj[att][self._convert_to_pascal_case(key)] = val.to_json()
else:
json_obj[att][key] = val
else:
json_obj[att] = cur_val
return json_obj


class DeploymentConfig(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config."""

__slots__ = [
"model_data_download_timeout",
"container_startup_health_check_timeout",
"image_uri",
"model_data",
"instance_type",
"environment",
"compute_resource_requirements",
]

def __init__(
self, init_kwargs: JumpStartModelInitKwargs, deploy_kwargs: JumpStartModelDeployKwargs
):
"""Instantiates DeploymentConfig object."""
if init_kwargs is not None:
self.image_uri = init_kwargs.image_uri
self.model_data = init_kwargs.model_data
self.instance_type = init_kwargs.instance_type
self.environment = init_kwargs.env
if init_kwargs.resources is not None:
self.compute_resource_requirements = (
init_kwargs.resources.get_compute_resource_requirements()
)
if deploy_kwargs is not None:
self.model_data_download_timeout = deploy_kwargs.model_data_download_timeout
self.container_startup_health_check_timeout = (
deploy_kwargs.container_startup_health_check_timeout
)


class DeploymentConfigMetadata(BaseDeploymentConfigDataHolder):
"""Dataclass representing a Deployment Config Metadata"""

__slots__ = [
"config_name",
"benchmark_metrics",
"deployment_config",
]

def __init__(
self,
config_name: str,
benchmark_metrics: List[JumpStartBenchmarkStat],
init_kwargs: JumpStartModelInitKwargs,
deploy_kwargs: JumpStartModelDeployKwargs,
):
"""Instantiates DeploymentConfigMetadata object."""
self.config_name = config_name
self.benchmark_metrics = benchmark_metrics
self.deployment_config = DeploymentConfig(init_kwargs, deploy_kwargs)
41 changes: 41 additions & 0 deletions src/sagemaker/jumpstart/utils.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -999,3 +999,44 @@ def get_jumpstart_configs(
if metadata_configs
else {}
)


def extract_metrics_from_deployment_configs(
deployment_configs: List[Dict[str, Any]], config_name: str
) -> Dict[str, List[str]]:
"""Extracts metrics from deployment configs.

Args:
deployment_configs (list[dict[str, Any]]): List of deployment configs.
config_name (str): The name of the deployment config use by the model.
"""

data = {"Config Name": [], "Instance Type": [], "Selected": []}

for index, deployment_config in enumerate(deployment_configs):
if deployment_config.get("DeploymentConfig") is None:
continue

benchmark_metrics = deployment_config.get("BenchmarkMetrics")
if benchmark_metrics is not None:
data["Config Name"].append(deployment_config.get("ConfigName"))
data["Instance Type"].append(
deployment_config.get("DeploymentConfig").get("InstanceType")
Comment thread
makungaj1 marked this conversation as resolved.
)
data["Selected"].append(
"Yes"
if (config_name is not None and config_name == deployment_config.get("ConfigName"))
else "No"
)

if index == 0:
for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
data[column_name] = []

for benchmark_metric in benchmark_metrics:
column_name = f"{benchmark_metric.get('name')} ({benchmark_metric.get('unit')})"
if column_name in data.keys():
data[column_name].append(benchmark_metric.get("value"))

return data
14 changes: 13 additions & 1 deletion src/sagemaker/serve/builder/jumpstart_builder.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -16,7 +16,7 @@
import copy
from abc import ABC, abstractmethod
from datetime import datetime, timedelta
from typing import Type
from typing import Type, Any, List, Dict
import logging

from sagemaker.model import Model
Expand DownExpand Up@@ -431,6 +431,18 @@ def tune_for_tgi_jumpstart(self, max_tuning_duration: int = 1800):
sharded_supported=sharded_supported, max_tuning_duration=max_tuning_duration
)

def display_benchmark_metrics(self):
"""Display Markdown Benchmark Metrics for deployment configs."""
self.pysdk_model.display_benchmark_metrics()

def list_deployment_configs(self) -> List[Dict[str, Any]]:
"""List deployment configs for ``This`` model in the current region.

Returns:
List[Dict[str, Any]]: A list of deployment configs.
"""
return self.pysdk_model.list_deployment_configs()

def _build_for_jumpstart(self):
"""Placeholder docstring"""
# we do not pickle for jumpstart. set to none
Expand Down
Loading