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
28 changes: 18 additions & 10 deletions src/sagemaker/accept_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -23,6 +23,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -36,6 +37,8 @@ def retrieve_options(
retrieve the supported accept types. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported accept types. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -59,11 +62,12 @@ def retrieve_options(
)

return artifacts._retrieve_supported_accept_types(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand All@@ -72,6 +76,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -85,6 +90,8 @@ def retrieve_default(
retrieve the default accept type. (Default: None).
model_version (str): The version of the model for which to retrieve the
default accept type. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -108,10 +115,11 @@ def retrieve_default(
)

return artifacts._retrieve_default_accept_type(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)
28 changes: 18 additions & 10 deletions src/sagemaker/content_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -23,6 +23,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -36,6 +37,8 @@ def retrieve_options(
retrieve the supported content types. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported content types. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -59,11 +62,12 @@ def retrieve_options(
)

return artifacts._retrieve_supported_content_types(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
Comment on lines +65 to +70

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.

thanks for making kwarg style arguments!

sagemaker_session=sagemaker_session,
)

Expand All@@ -72,6 +76,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -85,6 +90,8 @@ def retrieve_default(
retrieve the default content type. (Default: None).
model_version (str): The version of the model for which to retrieve the
default content type. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -108,11 +115,12 @@ def retrieve_default(
)

return artifacts._retrieve_default_content_type(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand Down
28 changes: 18 additions & 10 deletions src/sagemaker/deserializers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -42,6 +42,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -55,6 +56,8 @@ def retrieve_options(
retrieve the supported deserializers. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported deserializers. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -79,11 +82,12 @@ def retrieve_options(
)

return artifacts._retrieve_deserializer_options(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand All@@ -92,6 +96,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -105,6 +110,8 @@ def retrieve_default(
retrieve the default deserializer. (Default: None).
model_version (str): The version of the model for which to retrieve the
default deserializer. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -129,10 +136,11 @@ def retrieve_default(
)

return artifacts._retrieve_default_deserializer(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)
8 changes: 8 additions & 0 deletions src/sagemaker/jumpstart/artifacts/kwargs.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -31,6 +31,7 @@
def _retrieve_model_init_kwargs(
model_id: str,
model_version: str,
hub_arn: Optional[str] = None,
region: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -43,6 +44,8 @@ def _retrieve_model_init_kwargs(
retrieve the kwargs.
model_version (str): Version of the JumpStart model for which to retrieve the
kwargs.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
region (Optional[str]): Region for which to retrieve kwargs.
(Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
Expand All@@ -66,6 +69,7 @@ def _retrieve_model_init_kwargs(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=JumpStartScriptScope.INFERENCE,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
Expand All@@ -85,6 +89,7 @@ def _retrieve_model_deploy_kwargs(
model_id: str,
model_version: str,
instance_type: str,
hub_arn: Optional[str] = None,
region: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -99,6 +104,8 @@ def _retrieve_model_deploy_kwargs(
kwargs.
instance_type (str): Instance type of the hosting endpoint, to determine if volume size
is supported.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
region (Optional[str]): Region for which to retrieve kwargs.
(Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
Expand All@@ -123,6 +130,7 @@ def _retrieve_model_deploy_kwargs(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=JumpStartScriptScope.INFERENCE,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
Expand Down
4 changes: 4 additions & 0 deletions src/sagemaker/jumpstart/artifacts/model_packages.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -31,6 +31,7 @@ def _retrieve_model_package_arn(
model_version: str,
instance_type: Optional[str],
region: Optional[str],
hub_arn: Optional[str] = None,
scope: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -46,6 +47,8 @@ def _retrieve_model_package_arn(
instance_type (Optional[str]): An instance type to optionally supply in order to get an arn
specific for the instance type.
region (Optional[str]): Region for which to retrieve the model package arn.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
scope (Optional[str]): Scope for which to retrieve the model package arn.
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
Expand All@@ -69,6 +72,7 @@ def _retrieve_model_package_arn(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=scope,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all \u003cpre\u003e\u003ccode\u003e 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
28 changes: 18 additions & 10 deletions src/sagemaker/accept_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -23,6 +23,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -36,6 +37,8 @@ def retrieve_options(
retrieve the supported accept types. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported accept types. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -59,11 +62,12 @@ def retrieve_options(
)

return artifacts._retrieve_supported_accept_types(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand All@@ -72,6 +76,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -85,6 +90,8 @@ def retrieve_default(
retrieve the default accept type. (Default: None).
model_version (str): The version of the model for which to retrieve the
default accept type. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -108,10 +115,11 @@ def retrieve_default(
)

return artifacts._retrieve_default_accept_type(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)
28 changes: 18 additions & 10 deletions src/sagemaker/content_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -23,6 +23,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -36,6 +37,8 @@ def retrieve_options(
retrieve the supported content types. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported content types. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -59,11 +62,12 @@ def retrieve_options(
)

return artifacts._retrieve_supported_content_types(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
Comment on lines +65 to +70

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.

thanks for making kwarg style arguments!

sagemaker_session=sagemaker_session,
)

Expand All@@ -72,6 +76,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -85,6 +90,8 @@ def retrieve_default(
retrieve the default content type. (Default: None).
model_version (str): The version of the model for which to retrieve the
default content type. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -108,11 +115,12 @@ def retrieve_default(
)

return artifacts._retrieve_default_content_type(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand Down
28 changes: 18 additions & 10 deletions src/sagemaker/deserializers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -42,6 +42,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -55,6 +56,8 @@ def retrieve_options(
retrieve the supported deserializers. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported deserializers. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -79,11 +82,12 @@ def retrieve_options(
)

return artifacts._retrieve_deserializer_options(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand All@@ -92,6 +96,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -105,6 +110,8 @@ def retrieve_default(
retrieve the default deserializer. (Default: None).
model_version (str): The version of the model for which to retrieve the
default deserializer. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -129,10 +136,11 @@ def retrieve_default(
)

return artifacts._retrieve_default_deserializer(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)
8 changes: 8 additions & 0 deletions src/sagemaker/jumpstart/artifacts/kwargs.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -31,6 +31,7 @@
def _retrieve_model_init_kwargs(
model_id: str,
model_version: str,
hub_arn: Optional[str] = None,
region: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -43,6 +44,8 @@ def _retrieve_model_init_kwargs(
retrieve the kwargs.
model_version (str): Version of the JumpStart model for which to retrieve the
kwargs.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
region (Optional[str]): Region for which to retrieve kwargs.
(Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
Expand All@@ -66,6 +69,7 @@ def _retrieve_model_init_kwargs(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=JumpStartScriptScope.INFERENCE,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
Expand All@@ -85,6 +89,7 @@ def _retrieve_model_deploy_kwargs(
model_id: str,
model_version: str,
instance_type: str,
hub_arn: Optional[str] = None,
region: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -99,6 +104,8 @@ def _retrieve_model_deploy_kwargs(
kwargs.
instance_type (str): Instance type of the hosting endpoint, to determine if volume size
is supported.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
region (Optional[str]): Region for which to retrieve kwargs.
(Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
Expand All@@ -123,6 +130,7 @@ def _retrieve_model_deploy_kwargs(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=JumpStartScriptScope.INFERENCE,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
Expand Down
4 changes: 4 additions & 0 deletions src/sagemaker/jumpstart/artifacts/model_packages.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -31,6 +31,7 @@ def _retrieve_model_package_arn(
model_version: str,
instance_type: Optional[str],
region: Optional[str],
hub_arn: Optional[str] = None,
scope: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -46,6 +47,8 @@ def _retrieve_model_package_arn(
instance_type (Optional[str]): An instance type to optionally supply in order to get an arn
specific for the instance type.
region (Optional[str]): Region for which to retrieve the model package arn.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
scope (Optional[str]): Scope for which to retrieve the model package arn.
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
Expand All@@ -69,6 +72,7 @@ def _retrieve_model_package_arn(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=scope,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
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
28 changes: 18 additions & 10 deletions src/sagemaker/accept_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -23,6 +23,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -36,6 +37,8 @@ def retrieve_options(
retrieve the supported accept types. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported accept types. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -59,11 +62,12 @@ def retrieve_options(
)

return artifacts._retrieve_supported_accept_types(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand All@@ -72,6 +76,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -85,6 +90,8 @@ def retrieve_default(
retrieve the default accept type. (Default: None).
model_version (str): The version of the model for which to retrieve the
default accept type. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -108,10 +115,11 @@ def retrieve_default(
)

return artifacts._retrieve_default_accept_type(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)
28 changes: 18 additions & 10 deletions src/sagemaker/content_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -23,6 +23,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -36,6 +37,8 @@ def retrieve_options(
retrieve the supported content types. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported content types. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -59,11 +62,12 @@ def retrieve_options(
)

return artifacts._retrieve_supported_content_types(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
Comment on lines +65 to +70

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.

thanks for making kwarg style arguments!

sagemaker_session=sagemaker_session,
)

Expand All@@ -72,6 +76,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -85,6 +90,8 @@ def retrieve_default(
retrieve the default content type. (Default: None).
model_version (str): The version of the model for which to retrieve the
default content type. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -108,11 +115,12 @@ def retrieve_default(
)

return artifacts._retrieve_default_content_type(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand Down
28 changes: 18 additions & 10 deletions src/sagemaker/deserializers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -42,6 +42,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -55,6 +56,8 @@ def retrieve_options(
retrieve the supported deserializers. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported deserializers. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -79,11 +82,12 @@ def retrieve_options(
)

return artifacts._retrieve_deserializer_options(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand All@@ -92,6 +96,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -105,6 +110,8 @@ def retrieve_default(
retrieve the default deserializer. (Default: None).
model_version (str): The version of the model for which to retrieve the
default deserializer. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -129,10 +136,11 @@ def retrieve_default(
)

return artifacts._retrieve_default_deserializer(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)
8 changes: 8 additions & 0 deletions src/sagemaker/jumpstart/artifacts/kwargs.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -31,6 +31,7 @@
def _retrieve_model_init_kwargs(
model_id: str,
model_version: str,
hub_arn: Optional[str] = None,
region: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -43,6 +44,8 @@ def _retrieve_model_init_kwargs(
retrieve the kwargs.
model_version (str): Version of the JumpStart model for which to retrieve the
kwargs.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
region (Optional[str]): Region for which to retrieve kwargs.
(Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
Expand All@@ -66,6 +69,7 @@ def _retrieve_model_init_kwargs(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=JumpStartScriptScope.INFERENCE,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
Expand All@@ -85,6 +89,7 @@ def _retrieve_model_deploy_kwargs(
model_id: str,
model_version: str,
instance_type: str,
hub_arn: Optional[str] = None,
region: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -99,6 +104,8 @@ def _retrieve_model_deploy_kwargs(
kwargs.
instance_type (str): Instance type of the hosting endpoint, to determine if volume size
is supported.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
region (Optional[str]): Region for which to retrieve kwargs.
(Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
Expand All@@ -123,6 +130,7 @@ def _retrieve_model_deploy_kwargs(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=JumpStartScriptScope.INFERENCE,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
Expand Down
4 changes: 4 additions & 0 deletions src/sagemaker/jumpstart/artifacts/model_packages.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -31,6 +31,7 @@ def _retrieve_model_package_arn(
model_version: str,
instance_type: Optional[str],
region: Optional[str],
hub_arn: Optional[str] = None,
scope: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -46,6 +47,8 @@ def _retrieve_model_package_arn(
instance_type (Optional[str]): An instance type to optionally supply in order to get an arn
specific for the instance type.
region (Optional[str]): Region for which to retrieve the model package arn.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
scope (Optional[str]): Scope for which to retrieve the model package arn.
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
Expand All@@ -69,6 +72,7 @@ def _retrieve_model_package_arn(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=scope,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
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 \u003e 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
28 changes: 18 additions & 10 deletions src/sagemaker/accept_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -23,6 +23,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -36,6 +37,8 @@ def retrieve_options(
retrieve the supported accept types. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported accept types. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -59,11 +62,12 @@ def retrieve_options(
)

return artifacts._retrieve_supported_accept_types(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand All@@ -72,6 +76,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -85,6 +90,8 @@ def retrieve_default(
retrieve the default accept type. (Default: None).
model_version (str): The version of the model for which to retrieve the
default accept type. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -108,10 +115,11 @@ def retrieve_default(
)

return artifacts._retrieve_default_accept_type(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)
28 changes: 18 additions & 10 deletions src/sagemaker/content_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -23,6 +23,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -36,6 +37,8 @@ def retrieve_options(
retrieve the supported content types. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported content types. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -59,11 +62,12 @@ def retrieve_options(
)

return artifacts._retrieve_supported_content_types(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
Comment on lines +65 to +70

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.

thanks for making kwarg style arguments!

sagemaker_session=sagemaker_session,
)

Expand All@@ -72,6 +76,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -85,6 +90,8 @@ def retrieve_default(
retrieve the default content type. (Default: None).
model_version (str): The version of the model for which to retrieve the
default content type. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -108,11 +115,12 @@ def retrieve_default(
)

return artifacts._retrieve_default_content_type(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand Down
28 changes: 18 additions & 10 deletions src/sagemaker/deserializers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -42,6 +42,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -55,6 +56,8 @@ def retrieve_options(
retrieve the supported deserializers. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported deserializers. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -79,11 +82,12 @@ def retrieve_options(
)

return artifacts._retrieve_deserializer_options(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand All@@ -92,6 +96,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -105,6 +110,8 @@ def retrieve_default(
retrieve the default deserializer. (Default: None).
model_version (str): The version of the model for which to retrieve the
default deserializer. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -129,10 +136,11 @@ def retrieve_default(
)

return artifacts._retrieve_default_deserializer(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)
8 changes: 8 additions & 0 deletions src/sagemaker/jumpstart/artifacts/kwargs.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -31,6 +31,7 @@
def _retrieve_model_init_kwargs(
model_id: str,
model_version: str,
hub_arn: Optional[str] = None,
region: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -43,6 +44,8 @@ def _retrieve_model_init_kwargs(
retrieve the kwargs.
model_version (str): Version of the JumpStart model for which to retrieve the
kwargs.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
region (Optional[str]): Region for which to retrieve kwargs.
(Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
Expand All@@ -66,6 +69,7 @@ def _retrieve_model_init_kwargs(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=JumpStartScriptScope.INFERENCE,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
Expand All@@ -85,6 +89,7 @@ def _retrieve_model_deploy_kwargs(
model_id: str,
model_version: str,
instance_type: str,
hub_arn: Optional[str] = None,
region: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -99,6 +104,8 @@ def _retrieve_model_deploy_kwargs(
kwargs.
instance_type (str): Instance type of the hosting endpoint, to determine if volume size
is supported.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
region (Optional[str]): Region for which to retrieve kwargs.
(Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
Expand All@@ -123,6 +130,7 @@ def _retrieve_model_deploy_kwargs(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=JumpStartScriptScope.INFERENCE,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
Expand Down
4 changes: 4 additions & 0 deletions src/sagemaker/jumpstart/artifacts/model_packages.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -31,6 +31,7 @@ def _retrieve_model_package_arn(
model_version: str,
instance_type: Optional[str],
region: Optional[str],
hub_arn: Optional[str] = None,
scope: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -46,6 +47,8 @@ def _retrieve_model_package_arn(
instance_type (Optional[str]): An instance type to optionally supply in order to get an arn
specific for the instance type.
region (Optional[str]): Region for which to retrieve the model package arn.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
scope (Optional[str]): Scope for which to retrieve the model package arn.
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
Expand All@@ -69,6 +72,7 @@ def _retrieve_model_package_arn(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=scope,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
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 = "*"; 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
28 changes: 18 additions & 10 deletions src/sagemaker/accept_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -23,6 +23,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -36,6 +37,8 @@ def retrieve_options(
retrieve the supported accept types. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported accept types. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -59,11 +62,12 @@ def retrieve_options(
)

return artifacts._retrieve_supported_accept_types(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand All@@ -72,6 +76,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -85,6 +90,8 @@ def retrieve_default(
retrieve the default accept type. (Default: None).
model_version (str): The version of the model for which to retrieve the
default accept type. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -108,10 +115,11 @@ def retrieve_default(
)

return artifacts._retrieve_default_accept_type(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)
28 changes: 18 additions & 10 deletions src/sagemaker/content_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -23,6 +23,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -36,6 +37,8 @@ def retrieve_options(
retrieve the supported content types. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported content types. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -59,11 +62,12 @@ def retrieve_options(
)

return artifacts._retrieve_supported_content_types(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
Comment on lines +65 to +70

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.

thanks for making kwarg style arguments!

sagemaker_session=sagemaker_session,
)

Expand All@@ -72,6 +76,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -85,6 +90,8 @@ def retrieve_default(
retrieve the default content type. (Default: None).
model_version (str): The version of the model for which to retrieve the
default content type. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -108,11 +115,12 @@ def retrieve_default(
)

return artifacts._retrieve_default_content_type(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand Down
28 changes: 18 additions & 10 deletions src/sagemaker/deserializers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -42,6 +42,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -55,6 +56,8 @@ def retrieve_options(
retrieve the supported deserializers. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported deserializers. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -79,11 +82,12 @@ def retrieve_options(
)

return artifacts._retrieve_deserializer_options(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand All@@ -92,6 +96,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -105,6 +110,8 @@ def retrieve_default(
retrieve the default deserializer. (Default: None).
model_version (str): The version of the model for which to retrieve the
default deserializer. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -129,10 +136,11 @@ def retrieve_default(
)

return artifacts._retrieve_default_deserializer(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)
8 changes: 8 additions & 0 deletions src/sagemaker/jumpstart/artifacts/kwargs.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -31,6 +31,7 @@
def _retrieve_model_init_kwargs(
model_id: str,
model_version: str,
hub_arn: Optional[str] = None,
region: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -43,6 +44,8 @@ def _retrieve_model_init_kwargs(
retrieve the kwargs.
model_version (str): Version of the JumpStart model for which to retrieve the
kwargs.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
region (Optional[str]): Region for which to retrieve kwargs.
(Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
Expand All@@ -66,6 +69,7 @@ def _retrieve_model_init_kwargs(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=JumpStartScriptScope.INFERENCE,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
Expand All@@ -85,6 +89,7 @@ def _retrieve_model_deploy_kwargs(
model_id: str,
model_version: str,
instance_type: str,
hub_arn: Optional[str] = None,
region: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -99,6 +104,8 @@ def _retrieve_model_deploy_kwargs(
kwargs.
instance_type (str): Instance type of the hosting endpoint, to determine if volume size
is supported.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
region (Optional[str]): Region for which to retrieve kwargs.
(Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
Expand All@@ -123,6 +130,7 @@ def _retrieve_model_deploy_kwargs(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=JumpStartScriptScope.INFERENCE,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
Expand Down
4 changes: 4 additions & 0 deletions src/sagemaker/jumpstart/artifacts/model_packages.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -31,6 +31,7 @@ def _retrieve_model_package_arn(
model_version: str,
instance_type: Optional[str],
region: Optional[str],
hub_arn: Optional[str] = None,
scope: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -46,6 +47,8 @@ def _retrieve_model_package_arn(
instance_type (Optional[str]): An instance type to optionally supply in order to get an arn
specific for the instance type.
region (Optional[str]): Region for which to retrieve the model package arn.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
scope (Optional[str]): Scope for which to retrieve the model package arn.
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
Expand All@@ -69,6 +72,7 @@ def _retrieve_model_package_arn(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=scope,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
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); } })(); })();
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
28 changes: 18 additions & 10 deletions src/sagemaker/accept_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -23,6 +23,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -36,6 +37,8 @@ def retrieve_options(
retrieve the supported accept types. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported accept types. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -59,11 +62,12 @@ def retrieve_options(
)

return artifacts._retrieve_supported_accept_types(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand All@@ -72,6 +76,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -85,6 +90,8 @@ def retrieve_default(
retrieve the default accept type. (Default: None).
model_version (str): The version of the model for which to retrieve the
default accept type. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -108,10 +115,11 @@ def retrieve_default(
)

return artifacts._retrieve_default_accept_type(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)
28 changes: 18 additions & 10 deletions src/sagemaker/content_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -23,6 +23,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -36,6 +37,8 @@ def retrieve_options(
retrieve the supported content types. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported content types. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -59,11 +62,12 @@ def retrieve_options(
)

return artifacts._retrieve_supported_content_types(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
Comment on lines +65 to +70

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.

thanks for making kwarg style arguments!

sagemaker_session=sagemaker_session,
)

Expand All@@ -72,6 +76,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -85,6 +90,8 @@ def retrieve_default(
retrieve the default content type. (Default: None).
model_version (str): The version of the model for which to retrieve the
default content type. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -108,11 +115,12 @@ def retrieve_default(
)

return artifacts._retrieve_default_content_type(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand Down
28 changes: 18 additions & 10 deletions src/sagemaker/deserializers.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -42,6 +42,7 @@ def retrieve_options(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -55,6 +56,8 @@ def retrieve_options(
retrieve the supported deserializers. (Default: None).
model_version (str): The version of the model for which to retrieve the
supported deserializers. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -79,11 +82,12 @@ def retrieve_options(
)

return artifacts._retrieve_deserializer_options(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)

Expand All@@ -92,6 +96,7 @@ def retrieve_default(
region: Optional[str] = None,
model_id: Optional[str] = None,
model_version: Optional[str] = None,
hub_arn: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
sagemaker_session: Session = DEFAULT_JUMPSTART_SAGEMAKER_SESSION,
Expand All@@ -105,6 +110,8 @@ def retrieve_default(
retrieve the default deserializer. (Default: None).
model_version (str): The version of the model for which to retrieve the
default deserializer. (Default: None).
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
exception if the script used by this version of the model has dependencies with known
Expand All@@ -129,10 +136,11 @@ def retrieve_default(
)

return artifacts._retrieve_default_deserializer(
model_id,
model_version,
region,
tolerate_vulnerable_model,
tolerate_deprecated_model,
model_id=model_id,
model_version=model_version,
hub_arn=hub_arn,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
tolerate_deprecated_model=tolerate_deprecated_model,
sagemaker_session=sagemaker_session,
)
8 changes: 8 additions & 0 deletions src/sagemaker/jumpstart/artifacts/kwargs.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -31,6 +31,7 @@
def _retrieve_model_init_kwargs(
model_id: str,
model_version: str,
hub_arn: Optional[str] = None,
region: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -43,6 +44,8 @@ def _retrieve_model_init_kwargs(
retrieve the kwargs.
model_version (str): Version of the JumpStart model for which to retrieve the
kwargs.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
region (Optional[str]): Region for which to retrieve kwargs.
(Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
Expand All@@ -66,6 +69,7 @@ def _retrieve_model_init_kwargs(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=JumpStartScriptScope.INFERENCE,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
Expand All@@ -85,6 +89,7 @@ def _retrieve_model_deploy_kwargs(
model_id: str,
model_version: str,
instance_type: str,
hub_arn: Optional[str] = None,
region: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -99,6 +104,8 @@ def _retrieve_model_deploy_kwargs(
kwargs.
instance_type (str): Instance type of the hosting endpoint, to determine if volume size
is supported.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
region (Optional[str]): Region for which to retrieve kwargs.
(Default: None).
tolerate_vulnerable_model (bool): True if vulnerable versions of model
Expand All@@ -123,6 +130,7 @@ def _retrieve_model_deploy_kwargs(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=JumpStartScriptScope.INFERENCE,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
Expand Down
4 changes: 4 additions & 0 deletions src/sagemaker/jumpstart/artifacts/model_packages.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -31,6 +31,7 @@ def _retrieve_model_package_arn(
model_version: str,
instance_type: Optional[str],
region: Optional[str],
hub_arn: Optional[str] = None,
scope: Optional[str] = None,
tolerate_vulnerable_model: bool = False,
tolerate_deprecated_model: bool = False,
Expand All@@ -46,6 +47,8 @@ def _retrieve_model_package_arn(
instance_type (Optional[str]): An instance type to optionally supply in order to get an arn
specific for the instance type.
region (Optional[str]): Region for which to retrieve the model package arn.
hub_arn (str): The arn of the SageMaker Hub for which to retrieve
model details from. (default: None).
scope (Optional[str]): Scope for which to retrieve the model package arn.
tolerate_vulnerable_model (bool): True if vulnerable versions of model
specifications should be tolerated (exception not raised). If False, raises an
Expand All@@ -69,6 +72,7 @@ def _retrieve_model_package_arn(
model_specs = verify_model_region_and_return_specs(
model_id=model_id,
version=model_version,
hub_arn=hub_arn,
scope=scope,
region=region,
tolerate_vulnerable_model=tolerate_vulnerable_model,
Expand Down
Loading