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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions sagemaker-core/src/sagemaker/core/apiutils/_base_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -219,10 +219,29 @@ def with_boto(self, boto_dict):
)
return self

# Lineage entity creation methods whose requests are captured as pipeline
# step arguments when invoked under a ``PipelineSession`` (instead of
# calling the service). Used by ``sagemaker.mlops.workflow.LineageStep``.
_PIPELINE_CAPTURABLE_METHODS = frozenset(
{"create_action", "create_artifact", "create_context", "add_association"}
)

def _invoke_api(self, boto_method, boto_method_members):
"""Invoke a SageMaker API."""
api_values = {k: v for k, v in vars(self).items() if k in boto_method_members}
api_kwargs = self.to_boto(api_values)

if boto_method in self._PIPELINE_CAPTURABLE_METHODS:
# Lazy import to avoid a circular dependency at module load time.
from sagemaker.core.workflow.pipeline_context import (
PipelineSession,
_JobStepArguments,
)

if isinstance(self.sagemaker_session, PipelineSession):
self.sagemaker_session.context = _JobStepArguments(boto_method, api_kwargs)
return self.sagemaker_session.context

api_method = getattr(self.sagemaker_session.sagemaker_client, boto_method)
api_boto_response = api_method(**api_kwargs)
return self.with_boto(api_boto_response)
73 changes: 57 additions & 16 deletions sagemaker-core/src/sagemaker/core/helper/session_helper.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1091,15 +1091,23 @@ def endpoint_from_production_variants(
config_options["ExecutionRoleArn"] = role

logger.info("Creating endpoint-config with name %s", name)
self.sagemaker_client.create_endpoint_config(**config_options)

return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,

def submit(request):
self.sagemaker_client.create_endpoint_config(**request)
return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,
)

result = self._intercept_create_request(
config_options, submit, self.endpoint_from_production_variants.__name__
)
if self._is_pipeline_context():
return self.context
return result

def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live_logging=False):
"""Create an Amazon SageMaker ``Endpoint`` according to the configuration in the request.
Expand DownExpand Up@@ -1129,16 +1137,28 @@ def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live
tags = self._append_sagemaker_config_tags(
tags, "{}.{}.{}".format(SAGEMAKER, ENDPOINT, TAGS)
)
try:
res = self.sagemaker_client.create_endpoint(
EndpointName=endpoint_name, EndpointConfigName=config_name, Tags=tags
)
create_endpoint_request = {
"EndpointName": endpoint_name,
"EndpointConfigName": config_name,
"Tags": tags,
}

def submit(request):
res = self.sagemaker_client.create_endpoint(**request)
if res:
self.endpoint_arn = res["EndpointArn"]

if wait:
self.wait_for_endpoint(endpoint_name, live_logging=live_logging)
return endpoint_name

try:
result = self._intercept_create_request(
create_endpoint_request, submit, self.create_endpoint.__name__
)
if self._is_pipeline_context():
return self.context
return result
except Exception as e:
troubleshooting = (
"https://docs.aws.amazon.com/sagemaker/latest/dg/"
Expand DownExpand Up@@ -1257,10 +1277,18 @@ def create_inference_component(
if tags and len(tags) != 0:
request["Tags"] = tags

self.sagemaker_client.create_inference_component(**request)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name
def submit(req):
self.sagemaker_client.create_inference_component(**req)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name

result = self._intercept_create_request(
request, submit, self.create_inference_component.__name__
)
if self._is_pipeline_context():
return self.context
return result

def wait_for_inference_component(self, inference_component_name, poll=20):
"""Wait for an Amazon SageMaker ``Inference Component`` deployment to complete.
Expand DownExpand Up@@ -1439,6 +1467,19 @@ def _intercept_create_request(
"""
return create(request)

def _is_pipeline_context(self) -> bool:
"""Whether this session is a pipeline session capturing requests.

Producer methods that support composing pipeline steps use this to
return the captured step arguments (``self.context``) instead of the
result of a service call. Always ``False`` for a plain ``Session``.
"""
# Lazy import to avoid a circular dependency: pipeline_context imports
# from this module at import time.
from sagemaker.core.workflow.pipeline_context import PipelineSession

return isinstance(self, PipelineSession)

def _create_inference_recommendations_job_request(
self,
role: str,
Expand Down
8 changes: 8 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/__init__.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,7 @@
functions, conditions, properties) and can import from sagemaker.train and sagemaker.serve
for orchestration purposes.
"""

from __future__ import absolute_import

__version__ = "0.1.0"
Expand DownExpand Up@@ -46,8 +47,11 @@
from sagemaker.mlops.workflow.clarify_check_step import ClarifyCheckStep
from sagemaker.mlops.workflow.condition_step import ConditionStep
from sagemaker.mlops.workflow.emr_step import EMRStep, EMRStepConfig
from sagemaker.mlops.workflow.endpoint_step import EndpointConfigStep, EndpointStep
from sagemaker.mlops.workflow.fail_step import FailStep
from sagemaker.mlops.workflow.inference_component_step import InferenceComponentStep
from sagemaker.mlops.workflow.lambda_step import LambdaStep, LambdaOutput
from sagemaker.mlops.workflow.lineage_step import LineageStep
from sagemaker.mlops.workflow.model_step import ModelStep
from sagemaker.mlops.workflow.monitor_batch_transform_step import MonitorBatchTransformStep
from sagemaker.mlops.workflow.notebook_job_step import NotebookJobStep
Expand DownExpand Up@@ -98,9 +102,13 @@
"ConditionStep",
"EMRStep",
"EMRStepConfig",
"EndpointConfigStep",
"EndpointStep",
"FailStep",
"InferenceComponentStep",
"LambdaStep",
"LambdaOutput",
"LineageStep",
"ModelStep",
"MonitorBatchTransformStep",
"NotebookJobStep",
Expand Down
203 changes: 203 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/endpoint_step.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,203 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Step definitions for SageMaker Endpoint deployment in Pipelines.

These steps follow the ``step_args`` convention used by ``TrainingStep``
and ``ModelStep``: call the corresponding session method under a
:class:`~sagemaker.core.workflow.pipeline_context.PipelineSession` and
pass the returned step arguments to the step. The request is captured at
call time and the service call is deferred to pipeline execution.

Example::

pipeline_session = PipelineSession()

config_step_args = pipeline_session.endpoint_from_production_variants(
name="my-endpoint-config",
production_variants=[...],
)
config_step = EndpointConfigStep(name="CreateConfig", step_args=config_step_args)

endpoint_step_args = pipeline_session.create_endpoint(
endpoint_name="my-endpoint",
config_name="my-endpoint-config",
)
endpoint_step = EndpointStep(name="CreateEndpoint", step_args=endpoint_step_args)
"""

from __future__ import absolute_import

from typing import List, Optional, Union

from sagemaker.core.helper.pipeline_variable import RequestType
from sagemaker.core.workflow.pipeline_context import _JobStepArguments
from sagemaker.core.workflow.properties import Properties
from sagemaker.core.workflow.utilities import validate_step_args_input

from sagemaker.mlops.workflow.retry import RetryPolicy
from sagemaker.mlops.workflow.step_collections import StepCollection
from sagemaker.mlops.workflow.steps import (
CacheConfig,
ConfigurableRetryStep,
Step,
StepTypeEnum,
)


class EndpointConfigStep(ConfigurableRetryStep):
"""Creates a SageMaker EndpointConfig within a pipeline.

Wraps the SageMaker ``CreateEndpointConfig`` API. The ``step_args``
must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.endpoint_from_production_variants`
on a ``PipelineSession``.

``EndpointConfig`` is structurally cacheable (``cache_config``) and
retryable (``retry_policies``).
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
retry_policies: Optional[List[RetryPolicy]] = None,
):
"""Construct an ``EndpointConfigStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from
``pipeline_session.endpoint_from_production_variants()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
retry_policies (List[RetryPolicy]): Optional retry policies.
"""
super().__init__(
name=name,
step_type=StepTypeEnum.ENDPOINT_CONFIG,
display_name=display_name,
description=description,
depends_on=depends_on,
retry_policies=retry_policies,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"endpoint_from_production_variants"},
error_message=(
"The step_args of EndpointConfigStep must be obtained from "
"pipeline_session.endpoint_from_production_variants()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointConfigOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint_config``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointConfigOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict


class EndpointStep(Step):
"""Creates or updates a SageMaker Endpoint within a pipeline.

Wraps the SageMaker ``CreateEndpoint``/``UpdateEndpoint`` API -- the
pipeline chooses create-vs-update based on endpoint existence. The
``step_args`` must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.create_endpoint`
on a ``PipelineSession``.

``Endpoint`` is structurally cacheable but not retryable at the
pipeline level.
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
):
"""Construct an ``EndpointStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from ``pipeline_session.create_endpoint()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
"""
super().__init__(
name=name,
display_name=display_name,
description=description,
step_type=StepTypeEnum.ENDPOINT,
depends_on=depends_on,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"create_endpoint"},
error_message=(
"The step_args of EndpointStep must be obtained from "
"pipeline_session.create_endpoint()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict
Loading
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
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions sagemaker-core/src/sagemaker/core/apiutils/_base_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -219,10 +219,29 @@ def with_boto(self, boto_dict):
)
return self

# Lineage entity creation methods whose requests are captured as pipeline
# step arguments when invoked under a ``PipelineSession`` (instead of
# calling the service). Used by ``sagemaker.mlops.workflow.LineageStep``.
_PIPELINE_CAPTURABLE_METHODS = frozenset(
{"create_action", "create_artifact", "create_context", "add_association"}
)

def _invoke_api(self, boto_method, boto_method_members):
"""Invoke a SageMaker API."""
api_values = {k: v for k, v in vars(self).items() if k in boto_method_members}
api_kwargs = self.to_boto(api_values)

if boto_method in self._PIPELINE_CAPTURABLE_METHODS:
# Lazy import to avoid a circular dependency at module load time.
from sagemaker.core.workflow.pipeline_context import (
PipelineSession,
_JobStepArguments,
)

if isinstance(self.sagemaker_session, PipelineSession):
self.sagemaker_session.context = _JobStepArguments(boto_method, api_kwargs)
return self.sagemaker_session.context

api_method = getattr(self.sagemaker_session.sagemaker_client, boto_method)
api_boto_response = api_method(**api_kwargs)
return self.with_boto(api_boto_response)
73 changes: 57 additions & 16 deletions sagemaker-core/src/sagemaker/core/helper/session_helper.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1091,15 +1091,23 @@ def endpoint_from_production_variants(
config_options["ExecutionRoleArn"] = role

logger.info("Creating endpoint-config with name %s", name)
self.sagemaker_client.create_endpoint_config(**config_options)

return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,

def submit(request):
self.sagemaker_client.create_endpoint_config(**request)
return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,
)

result = self._intercept_create_request(
config_options, submit, self.endpoint_from_production_variants.__name__
)
if self._is_pipeline_context():
return self.context
return result

def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live_logging=False):
"""Create an Amazon SageMaker ``Endpoint`` according to the configuration in the request.
Expand DownExpand Up@@ -1129,16 +1137,28 @@ def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live
tags = self._append_sagemaker_config_tags(
tags, "{}.{}.{}".format(SAGEMAKER, ENDPOINT, TAGS)
)
try:
res = self.sagemaker_client.create_endpoint(
EndpointName=endpoint_name, EndpointConfigName=config_name, Tags=tags
)
create_endpoint_request = {
"EndpointName": endpoint_name,
"EndpointConfigName": config_name,
"Tags": tags,
}

def submit(request):
res = self.sagemaker_client.create_endpoint(**request)
if res:
self.endpoint_arn = res["EndpointArn"]

if wait:
self.wait_for_endpoint(endpoint_name, live_logging=live_logging)
return endpoint_name

try:
result = self._intercept_create_request(
create_endpoint_request, submit, self.create_endpoint.__name__
)
if self._is_pipeline_context():
return self.context
return result
except Exception as e:
troubleshooting = (
"https://docs.aws.amazon.com/sagemaker/latest/dg/"
Expand DownExpand Up@@ -1257,10 +1277,18 @@ def create_inference_component(
if tags and len(tags) != 0:
request["Tags"] = tags

self.sagemaker_client.create_inference_component(**request)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name
def submit(req):
self.sagemaker_client.create_inference_component(**req)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name

result = self._intercept_create_request(
request, submit, self.create_inference_component.__name__
)
if self._is_pipeline_context():
return self.context
return result

def wait_for_inference_component(self, inference_component_name, poll=20):
"""Wait for an Amazon SageMaker ``Inference Component`` deployment to complete.
Expand DownExpand Up@@ -1439,6 +1467,19 @@ def _intercept_create_request(
"""
return create(request)

def _is_pipeline_context(self) -> bool:
"""Whether this session is a pipeline session capturing requests.

Producer methods that support composing pipeline steps use this to
return the captured step arguments (``self.context``) instead of the
result of a service call. Always ``False`` for a plain ``Session``.
"""
# Lazy import to avoid a circular dependency: pipeline_context imports
# from this module at import time.
from sagemaker.core.workflow.pipeline_context import PipelineSession

return isinstance(self, PipelineSession)

def _create_inference_recommendations_job_request(
self,
role: str,
Expand Down
8 changes: 8 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/__init__.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,7 @@
functions, conditions, properties) and can import from sagemaker.train and sagemaker.serve
for orchestration purposes.
"""

from __future__ import absolute_import

__version__ = "0.1.0"
Expand DownExpand Up@@ -46,8 +47,11 @@
from sagemaker.mlops.workflow.clarify_check_step import ClarifyCheckStep
from sagemaker.mlops.workflow.condition_step import ConditionStep
from sagemaker.mlops.workflow.emr_step import EMRStep, EMRStepConfig
from sagemaker.mlops.workflow.endpoint_step import EndpointConfigStep, EndpointStep
from sagemaker.mlops.workflow.fail_step import FailStep
from sagemaker.mlops.workflow.inference_component_step import InferenceComponentStep
from sagemaker.mlops.workflow.lambda_step import LambdaStep, LambdaOutput
from sagemaker.mlops.workflow.lineage_step import LineageStep
from sagemaker.mlops.workflow.model_step import ModelStep
from sagemaker.mlops.workflow.monitor_batch_transform_step import MonitorBatchTransformStep
from sagemaker.mlops.workflow.notebook_job_step import NotebookJobStep
Expand DownExpand Up@@ -98,9 +102,13 @@
"ConditionStep",
"EMRStep",
"EMRStepConfig",
"EndpointConfigStep",
"EndpointStep",
"FailStep",
"InferenceComponentStep",
"LambdaStep",
"LambdaOutput",
"LineageStep",
"ModelStep",
"MonitorBatchTransformStep",
"NotebookJobStep",
Expand Down
203 changes: 203 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/endpoint_step.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,203 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Step definitions for SageMaker Endpoint deployment in Pipelines.

These steps follow the ``step_args`` convention used by ``TrainingStep``
and ``ModelStep``: call the corresponding session method under a
:class:`~sagemaker.core.workflow.pipeline_context.PipelineSession` and
pass the returned step arguments to the step. The request is captured at
call time and the service call is deferred to pipeline execution.

Example::

pipeline_session = PipelineSession()

config_step_args = pipeline_session.endpoint_from_production_variants(
name="my-endpoint-config",
production_variants=[...],
)
config_step = EndpointConfigStep(name="CreateConfig", step_args=config_step_args)

endpoint_step_args = pipeline_session.create_endpoint(
endpoint_name="my-endpoint",
config_name="my-endpoint-config",
)
endpoint_step = EndpointStep(name="CreateEndpoint", step_args=endpoint_step_args)
"""

from __future__ import absolute_import

from typing import List, Optional, Union

from sagemaker.core.helper.pipeline_variable import RequestType
from sagemaker.core.workflow.pipeline_context import _JobStepArguments
from sagemaker.core.workflow.properties import Properties
from sagemaker.core.workflow.utilities import validate_step_args_input

from sagemaker.mlops.workflow.retry import RetryPolicy
from sagemaker.mlops.workflow.step_collections import StepCollection
from sagemaker.mlops.workflow.steps import (
CacheConfig,
ConfigurableRetryStep,
Step,
StepTypeEnum,
)


class EndpointConfigStep(ConfigurableRetryStep):
"""Creates a SageMaker EndpointConfig within a pipeline.

Wraps the SageMaker ``CreateEndpointConfig`` API. The ``step_args``
must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.endpoint_from_production_variants`
on a ``PipelineSession``.

``EndpointConfig`` is structurally cacheable (``cache_config``) and
retryable (``retry_policies``).
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
retry_policies: Optional[List[RetryPolicy]] = None,
):
"""Construct an ``EndpointConfigStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from
``pipeline_session.endpoint_from_production_variants()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
retry_policies (List[RetryPolicy]): Optional retry policies.
"""
super().__init__(
name=name,
step_type=StepTypeEnum.ENDPOINT_CONFIG,
display_name=display_name,
description=description,
depends_on=depends_on,
retry_policies=retry_policies,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"endpoint_from_production_variants"},
error_message=(
"The step_args of EndpointConfigStep must be obtained from "
"pipeline_session.endpoint_from_production_variants()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointConfigOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint_config``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointConfigOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict


class EndpointStep(Step):
"""Creates or updates a SageMaker Endpoint within a pipeline.

Wraps the SageMaker ``CreateEndpoint``/``UpdateEndpoint`` API -- the
pipeline chooses create-vs-update based on endpoint existence. The
``step_args`` must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.create_endpoint`
on a ``PipelineSession``.

``Endpoint`` is structurally cacheable but not retryable at the
pipeline level.
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
):
"""Construct an ``EndpointStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from ``pipeline_session.create_endpoint()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
"""
super().__init__(
name=name,
display_name=display_name,
description=description,
step_type=StepTypeEnum.ENDPOINT,
depends_on=depends_on,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"create_endpoint"},
error_message=(
"The step_args of EndpointStep must be obtained from "
"pipeline_session.create_endpoint()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict
Loading
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
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions sagemaker-core/src/sagemaker/core/apiutils/_base_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -219,10 +219,29 @@ def with_boto(self, boto_dict):
)
return self

# Lineage entity creation methods whose requests are captured as pipeline
# step arguments when invoked under a ``PipelineSession`` (instead of
# calling the service). Used by ``sagemaker.mlops.workflow.LineageStep``.
_PIPELINE_CAPTURABLE_METHODS = frozenset(
{"create_action", "create_artifact", "create_context", "add_association"}
)

def _invoke_api(self, boto_method, boto_method_members):
"""Invoke a SageMaker API."""
api_values = {k: v for k, v in vars(self).items() if k in boto_method_members}
api_kwargs = self.to_boto(api_values)

if boto_method in self._PIPELINE_CAPTURABLE_METHODS:
# Lazy import to avoid a circular dependency at module load time.
from sagemaker.core.workflow.pipeline_context import (
PipelineSession,
_JobStepArguments,
)

if isinstance(self.sagemaker_session, PipelineSession):
self.sagemaker_session.context = _JobStepArguments(boto_method, api_kwargs)
return self.sagemaker_session.context

api_method = getattr(self.sagemaker_session.sagemaker_client, boto_method)
api_boto_response = api_method(**api_kwargs)
return self.with_boto(api_boto_response)
73 changes: 57 additions & 16 deletions sagemaker-core/src/sagemaker/core/helper/session_helper.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1091,15 +1091,23 @@ def endpoint_from_production_variants(
config_options["ExecutionRoleArn"] = role

logger.info("Creating endpoint-config with name %s", name)
self.sagemaker_client.create_endpoint_config(**config_options)

return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,

def submit(request):
self.sagemaker_client.create_endpoint_config(**request)
return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,
)

result = self._intercept_create_request(
config_options, submit, self.endpoint_from_production_variants.__name__
)
if self._is_pipeline_context():
return self.context
return result

def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live_logging=False):
"""Create an Amazon SageMaker ``Endpoint`` according to the configuration in the request.
Expand DownExpand Up@@ -1129,16 +1137,28 @@ def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live
tags = self._append_sagemaker_config_tags(
tags, "{}.{}.{}".format(SAGEMAKER, ENDPOINT, TAGS)
)
try:
res = self.sagemaker_client.create_endpoint(
EndpointName=endpoint_name, EndpointConfigName=config_name, Tags=tags
)
create_endpoint_request = {
"EndpointName": endpoint_name,
"EndpointConfigName": config_name,
"Tags": tags,
}

def submit(request):
res = self.sagemaker_client.create_endpoint(**request)
if res:
self.endpoint_arn = res["EndpointArn"]

if wait:
self.wait_for_endpoint(endpoint_name, live_logging=live_logging)
return endpoint_name

try:
result = self._intercept_create_request(
create_endpoint_request, submit, self.create_endpoint.__name__
)
if self._is_pipeline_context():
return self.context
return result
except Exception as e:
troubleshooting = (
"https://docs.aws.amazon.com/sagemaker/latest/dg/"
Expand DownExpand Up@@ -1257,10 +1277,18 @@ def create_inference_component(
if tags and len(tags) != 0:
request["Tags"] = tags

self.sagemaker_client.create_inference_component(**request)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name
def submit(req):
self.sagemaker_client.create_inference_component(**req)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name

result = self._intercept_create_request(
request, submit, self.create_inference_component.__name__
)
if self._is_pipeline_context():
return self.context
return result

def wait_for_inference_component(self, inference_component_name, poll=20):
"""Wait for an Amazon SageMaker ``Inference Component`` deployment to complete.
Expand DownExpand Up@@ -1439,6 +1467,19 @@ def _intercept_create_request(
"""
return create(request)

def _is_pipeline_context(self) -> bool:
"""Whether this session is a pipeline session capturing requests.

Producer methods that support composing pipeline steps use this to
return the captured step arguments (``self.context``) instead of the
result of a service call. Always ``False`` for a plain ``Session``.
"""
# Lazy import to avoid a circular dependency: pipeline_context imports
# from this module at import time.
from sagemaker.core.workflow.pipeline_context import PipelineSession

return isinstance(self, PipelineSession)

def _create_inference_recommendations_job_request(
self,
role: str,
Expand Down
8 changes: 8 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/__init__.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,7 @@
functions, conditions, properties) and can import from sagemaker.train and sagemaker.serve
for orchestration purposes.
"""

from __future__ import absolute_import

__version__ = "0.1.0"
Expand DownExpand Up@@ -46,8 +47,11 @@
from sagemaker.mlops.workflow.clarify_check_step import ClarifyCheckStep
from sagemaker.mlops.workflow.condition_step import ConditionStep
from sagemaker.mlops.workflow.emr_step import EMRStep, EMRStepConfig
from sagemaker.mlops.workflow.endpoint_step import EndpointConfigStep, EndpointStep
from sagemaker.mlops.workflow.fail_step import FailStep
from sagemaker.mlops.workflow.inference_component_step import InferenceComponentStep
from sagemaker.mlops.workflow.lambda_step import LambdaStep, LambdaOutput
from sagemaker.mlops.workflow.lineage_step import LineageStep
from sagemaker.mlops.workflow.model_step import ModelStep
from sagemaker.mlops.workflow.monitor_batch_transform_step import MonitorBatchTransformStep
from sagemaker.mlops.workflow.notebook_job_step import NotebookJobStep
Expand DownExpand Up@@ -98,9 +102,13 @@
"ConditionStep",
"EMRStep",
"EMRStepConfig",
"EndpointConfigStep",
"EndpointStep",
"FailStep",
"InferenceComponentStep",
"LambdaStep",
"LambdaOutput",
"LineageStep",
"ModelStep",
"MonitorBatchTransformStep",
"NotebookJobStep",
Expand Down
203 changes: 203 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/endpoint_step.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,203 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Step definitions for SageMaker Endpoint deployment in Pipelines.

These steps follow the ``step_args`` convention used by ``TrainingStep``
and ``ModelStep``: call the corresponding session method under a
:class:`~sagemaker.core.workflow.pipeline_context.PipelineSession` and
pass the returned step arguments to the step. The request is captured at
call time and the service call is deferred to pipeline execution.

Example::

pipeline_session = PipelineSession()

config_step_args = pipeline_session.endpoint_from_production_variants(
name="my-endpoint-config",
production_variants=[...],
)
config_step = EndpointConfigStep(name="CreateConfig", step_args=config_step_args)

endpoint_step_args = pipeline_session.create_endpoint(
endpoint_name="my-endpoint",
config_name="my-endpoint-config",
)
endpoint_step = EndpointStep(name="CreateEndpoint", step_args=endpoint_step_args)
"""

from __future__ import absolute_import

from typing import List, Optional, Union

from sagemaker.core.helper.pipeline_variable import RequestType
from sagemaker.core.workflow.pipeline_context import _JobStepArguments
from sagemaker.core.workflow.properties import Properties
from sagemaker.core.workflow.utilities import validate_step_args_input

from sagemaker.mlops.workflow.retry import RetryPolicy
from sagemaker.mlops.workflow.step_collections import StepCollection
from sagemaker.mlops.workflow.steps import (
CacheConfig,
ConfigurableRetryStep,
Step,
StepTypeEnum,
)


class EndpointConfigStep(ConfigurableRetryStep):
"""Creates a SageMaker EndpointConfig within a pipeline.

Wraps the SageMaker ``CreateEndpointConfig`` API. The ``step_args``
must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.endpoint_from_production_variants`
on a ``PipelineSession``.

``EndpointConfig`` is structurally cacheable (``cache_config``) and
retryable (``retry_policies``).
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
retry_policies: Optional[List[RetryPolicy]] = None,
):
"""Construct an ``EndpointConfigStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from
``pipeline_session.endpoint_from_production_variants()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
retry_policies (List[RetryPolicy]): Optional retry policies.
"""
super().__init__(
name=name,
step_type=StepTypeEnum.ENDPOINT_CONFIG,
display_name=display_name,
description=description,
depends_on=depends_on,
retry_policies=retry_policies,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"endpoint_from_production_variants"},
error_message=(
"The step_args of EndpointConfigStep must be obtained from "
"pipeline_session.endpoint_from_production_variants()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointConfigOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint_config``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointConfigOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict


class EndpointStep(Step):
"""Creates or updates a SageMaker Endpoint within a pipeline.

Wraps the SageMaker ``CreateEndpoint``/``UpdateEndpoint`` API -- the
pipeline chooses create-vs-update based on endpoint existence. The
``step_args`` must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.create_endpoint`
on a ``PipelineSession``.

``Endpoint`` is structurally cacheable but not retryable at the
pipeline level.
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
):
"""Construct an ``EndpointStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from ``pipeline_session.create_endpoint()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
"""
super().__init__(
name=name,
display_name=display_name,
description=description,
step_type=StepTypeEnum.ENDPOINT,
depends_on=depends_on,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"create_endpoint"},
error_message=(
"The step_args of EndpointStep must be obtained from "
"pipeline_session.create_endpoint()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict
Loading
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
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions sagemaker-core/src/sagemaker/core/apiutils/_base_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -219,10 +219,29 @@ def with_boto(self, boto_dict):
)
return self

# Lineage entity creation methods whose requests are captured as pipeline
# step arguments when invoked under a ``PipelineSession`` (instead of
# calling the service). Used by ``sagemaker.mlops.workflow.LineageStep``.
_PIPELINE_CAPTURABLE_METHODS = frozenset(
{"create_action", "create_artifact", "create_context", "add_association"}
)

def _invoke_api(self, boto_method, boto_method_members):
"""Invoke a SageMaker API."""
api_values = {k: v for k, v in vars(self).items() if k in boto_method_members}
api_kwargs = self.to_boto(api_values)

if boto_method in self._PIPELINE_CAPTURABLE_METHODS:
# Lazy import to avoid a circular dependency at module load time.
from sagemaker.core.workflow.pipeline_context import (
PipelineSession,
_JobStepArguments,
)

if isinstance(self.sagemaker_session, PipelineSession):
self.sagemaker_session.context = _JobStepArguments(boto_method, api_kwargs)
return self.sagemaker_session.context

api_method = getattr(self.sagemaker_session.sagemaker_client, boto_method)
api_boto_response = api_method(**api_kwargs)
return self.with_boto(api_boto_response)
73 changes: 57 additions & 16 deletions sagemaker-core/src/sagemaker/core/helper/session_helper.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1091,15 +1091,23 @@ def endpoint_from_production_variants(
config_options["ExecutionRoleArn"] = role

logger.info("Creating endpoint-config with name %s", name)
self.sagemaker_client.create_endpoint_config(**config_options)

return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,

def submit(request):
self.sagemaker_client.create_endpoint_config(**request)
return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,
)

result = self._intercept_create_request(
config_options, submit, self.endpoint_from_production_variants.__name__
)
if self._is_pipeline_context():
return self.context
return result

def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live_logging=False):
"""Create an Amazon SageMaker ``Endpoint`` according to the configuration in the request.
Expand DownExpand Up@@ -1129,16 +1137,28 @@ def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live
tags = self._append_sagemaker_config_tags(
tags, "{}.{}.{}".format(SAGEMAKER, ENDPOINT, TAGS)
)
try:
res = self.sagemaker_client.create_endpoint(
EndpointName=endpoint_name, EndpointConfigName=config_name, Tags=tags
)
create_endpoint_request = {
"EndpointName": endpoint_name,
"EndpointConfigName": config_name,
"Tags": tags,
}

def submit(request):
res = self.sagemaker_client.create_endpoint(**request)
if res:
self.endpoint_arn = res["EndpointArn"]

if wait:
self.wait_for_endpoint(endpoint_name, live_logging=live_logging)
return endpoint_name

try:
result = self._intercept_create_request(
create_endpoint_request, submit, self.create_endpoint.__name__
)
if self._is_pipeline_context():
return self.context
return result
except Exception as e:
troubleshooting = (
"https://docs.aws.amazon.com/sagemaker/latest/dg/"
Expand DownExpand Up@@ -1257,10 +1277,18 @@ def create_inference_component(
if tags and len(tags) != 0:
request["Tags"] = tags

self.sagemaker_client.create_inference_component(**request)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name
def submit(req):
self.sagemaker_client.create_inference_component(**req)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name

result = self._intercept_create_request(
request, submit, self.create_inference_component.__name__
)
if self._is_pipeline_context():
return self.context
return result

def wait_for_inference_component(self, inference_component_name, poll=20):
"""Wait for an Amazon SageMaker ``Inference Component`` deployment to complete.
Expand DownExpand Up@@ -1439,6 +1467,19 @@ def _intercept_create_request(
"""
return create(request)

def _is_pipeline_context(self) -> bool:
"""Whether this session is a pipeline session capturing requests.

Producer methods that support composing pipeline steps use this to
return the captured step arguments (``self.context``) instead of the
result of a service call. Always ``False`` for a plain ``Session``.
"""
# Lazy import to avoid a circular dependency: pipeline_context imports
# from this module at import time.
from sagemaker.core.workflow.pipeline_context import PipelineSession

return isinstance(self, PipelineSession)

def _create_inference_recommendations_job_request(
self,
role: str,
Expand Down
8 changes: 8 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/__init__.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,7 @@
functions, conditions, properties) and can import from sagemaker.train and sagemaker.serve
for orchestration purposes.
"""

from __future__ import absolute_import

__version__ = "0.1.0"
Expand DownExpand Up@@ -46,8 +47,11 @@
from sagemaker.mlops.workflow.clarify_check_step import ClarifyCheckStep
from sagemaker.mlops.workflow.condition_step import ConditionStep
from sagemaker.mlops.workflow.emr_step import EMRStep, EMRStepConfig
from sagemaker.mlops.workflow.endpoint_step import EndpointConfigStep, EndpointStep
from sagemaker.mlops.workflow.fail_step import FailStep
from sagemaker.mlops.workflow.inference_component_step import InferenceComponentStep
from sagemaker.mlops.workflow.lambda_step import LambdaStep, LambdaOutput
from sagemaker.mlops.workflow.lineage_step import LineageStep
from sagemaker.mlops.workflow.model_step import ModelStep
from sagemaker.mlops.workflow.monitor_batch_transform_step import MonitorBatchTransformStep
from sagemaker.mlops.workflow.notebook_job_step import NotebookJobStep
Expand DownExpand Up@@ -98,9 +102,13 @@
"ConditionStep",
"EMRStep",
"EMRStepConfig",
"EndpointConfigStep",
"EndpointStep",
"FailStep",
"InferenceComponentStep",
"LambdaStep",
"LambdaOutput",
"LineageStep",
"ModelStep",
"MonitorBatchTransformStep",
"NotebookJobStep",
Expand Down
203 changes: 203 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/endpoint_step.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,203 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Step definitions for SageMaker Endpoint deployment in Pipelines.

These steps follow the ``step_args`` convention used by ``TrainingStep``
and ``ModelStep``: call the corresponding session method under a
:class:`~sagemaker.core.workflow.pipeline_context.PipelineSession` and
pass the returned step arguments to the step. The request is captured at
call time and the service call is deferred to pipeline execution.

Example::

pipeline_session = PipelineSession()

config_step_args = pipeline_session.endpoint_from_production_variants(
name="my-endpoint-config",
production_variants=[...],
)
config_step = EndpointConfigStep(name="CreateConfig", step_args=config_step_args)

endpoint_step_args = pipeline_session.create_endpoint(
endpoint_name="my-endpoint",
config_name="my-endpoint-config",
)
endpoint_step = EndpointStep(name="CreateEndpoint", step_args=endpoint_step_args)
"""

from __future__ import absolute_import

from typing import List, Optional, Union

from sagemaker.core.helper.pipeline_variable import RequestType
from sagemaker.core.workflow.pipeline_context import _JobStepArguments
from sagemaker.core.workflow.properties import Properties
from sagemaker.core.workflow.utilities import validate_step_args_input

from sagemaker.mlops.workflow.retry import RetryPolicy
from sagemaker.mlops.workflow.step_collections import StepCollection
from sagemaker.mlops.workflow.steps import (
CacheConfig,
ConfigurableRetryStep,
Step,
StepTypeEnum,
)


class EndpointConfigStep(ConfigurableRetryStep):
"""Creates a SageMaker EndpointConfig within a pipeline.

Wraps the SageMaker ``CreateEndpointConfig`` API. The ``step_args``
must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.endpoint_from_production_variants`
on a ``PipelineSession``.

``EndpointConfig`` is structurally cacheable (``cache_config``) and
retryable (``retry_policies``).
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
retry_policies: Optional[List[RetryPolicy]] = None,
):
"""Construct an ``EndpointConfigStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from
``pipeline_session.endpoint_from_production_variants()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
retry_policies (List[RetryPolicy]): Optional retry policies.
"""
super().__init__(
name=name,
step_type=StepTypeEnum.ENDPOINT_CONFIG,
display_name=display_name,
description=description,
depends_on=depends_on,
retry_policies=retry_policies,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"endpoint_from_production_variants"},
error_message=(
"The step_args of EndpointConfigStep must be obtained from "
"pipeline_session.endpoint_from_production_variants()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointConfigOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint_config``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointConfigOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict


class EndpointStep(Step):
"""Creates or updates a SageMaker Endpoint within a pipeline.

Wraps the SageMaker ``CreateEndpoint``/``UpdateEndpoint`` API -- the
pipeline chooses create-vs-update based on endpoint existence. The
``step_args`` must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.create_endpoint`
on a ``PipelineSession``.

``Endpoint`` is structurally cacheable but not retryable at the
pipeline level.
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
):
"""Construct an ``EndpointStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from ``pipeline_session.create_endpoint()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
"""
super().__init__(
name=name,
display_name=display_name,
description=description,
step_type=StepTypeEnum.ENDPOINT,
depends_on=depends_on,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"create_endpoint"},
error_message=(
"The step_args of EndpointStep must be obtained from "
"pipeline_session.create_endpoint()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict
Loading
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
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions sagemaker-core/src/sagemaker/core/apiutils/_base_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -219,10 +219,29 @@ def with_boto(self, boto_dict):
)
return self

# Lineage entity creation methods whose requests are captured as pipeline
# step arguments when invoked under a ``PipelineSession`` (instead of
# calling the service). Used by ``sagemaker.mlops.workflow.LineageStep``.
_PIPELINE_CAPTURABLE_METHODS = frozenset(
{"create_action", "create_artifact", "create_context", "add_association"}
)

def _invoke_api(self, boto_method, boto_method_members):
"""Invoke a SageMaker API."""
api_values = {k: v for k, v in vars(self).items() if k in boto_method_members}
api_kwargs = self.to_boto(api_values)

if boto_method in self._PIPELINE_CAPTURABLE_METHODS:
# Lazy import to avoid a circular dependency at module load time.
from sagemaker.core.workflow.pipeline_context import (
PipelineSession,
_JobStepArguments,
)

if isinstance(self.sagemaker_session, PipelineSession):
self.sagemaker_session.context = _JobStepArguments(boto_method, api_kwargs)
return self.sagemaker_session.context

api_method = getattr(self.sagemaker_session.sagemaker_client, boto_method)
api_boto_response = api_method(**api_kwargs)
return self.with_boto(api_boto_response)
73 changes: 57 additions & 16 deletions sagemaker-core/src/sagemaker/core/helper/session_helper.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1091,15 +1091,23 @@ def endpoint_from_production_variants(
config_options["ExecutionRoleArn"] = role

logger.info("Creating endpoint-config with name %s", name)
self.sagemaker_client.create_endpoint_config(**config_options)

return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,

def submit(request):
self.sagemaker_client.create_endpoint_config(**request)
return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,
)

result = self._intercept_create_request(
config_options, submit, self.endpoint_from_production_variants.__name__
)
if self._is_pipeline_context():
return self.context
return result

def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live_logging=False):
"""Create an Amazon SageMaker ``Endpoint`` according to the configuration in the request.
Expand DownExpand Up@@ -1129,16 +1137,28 @@ def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live
tags = self._append_sagemaker_config_tags(
tags, "{}.{}.{}".format(SAGEMAKER, ENDPOINT, TAGS)
)
try:
res = self.sagemaker_client.create_endpoint(
EndpointName=endpoint_name, EndpointConfigName=config_name, Tags=tags
)
create_endpoint_request = {
"EndpointName": endpoint_name,
"EndpointConfigName": config_name,
"Tags": tags,
}

def submit(request):
res = self.sagemaker_client.create_endpoint(**request)
if res:
self.endpoint_arn = res["EndpointArn"]

if wait:
self.wait_for_endpoint(endpoint_name, live_logging=live_logging)
return endpoint_name

try:
result = self._intercept_create_request(
create_endpoint_request, submit, self.create_endpoint.__name__
)
if self._is_pipeline_context():
return self.context
return result
except Exception as e:
troubleshooting = (
"https://docs.aws.amazon.com/sagemaker/latest/dg/"
Expand DownExpand Up@@ -1257,10 +1277,18 @@ def create_inference_component(
if tags and len(tags) != 0:
request["Tags"] = tags

self.sagemaker_client.create_inference_component(**request)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name
def submit(req):
self.sagemaker_client.create_inference_component(**req)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name

result = self._intercept_create_request(
request, submit, self.create_inference_component.__name__
)
if self._is_pipeline_context():
return self.context
return result

def wait_for_inference_component(self, inference_component_name, poll=20):
"""Wait for an Amazon SageMaker ``Inference Component`` deployment to complete.
Expand DownExpand Up@@ -1439,6 +1467,19 @@ def _intercept_create_request(
"""
return create(request)

def _is_pipeline_context(self) -> bool:
"""Whether this session is a pipeline session capturing requests.

Producer methods that support composing pipeline steps use this to
return the captured step arguments (``self.context``) instead of the
result of a service call. Always ``False`` for a plain ``Session``.
"""
# Lazy import to avoid a circular dependency: pipeline_context imports
# from this module at import time.
from sagemaker.core.workflow.pipeline_context import PipelineSession

return isinstance(self, PipelineSession)

def _create_inference_recommendations_job_request(
self,
role: str,
Expand Down
8 changes: 8 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/__init__.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,7 @@
functions, conditions, properties) and can import from sagemaker.train and sagemaker.serve
for orchestration purposes.
"""

from __future__ import absolute_import

__version__ = "0.1.0"
Expand DownExpand Up@@ -46,8 +47,11 @@
from sagemaker.mlops.workflow.clarify_check_step import ClarifyCheckStep
from sagemaker.mlops.workflow.condition_step import ConditionStep
from sagemaker.mlops.workflow.emr_step import EMRStep, EMRStepConfig
from sagemaker.mlops.workflow.endpoint_step import EndpointConfigStep, EndpointStep
from sagemaker.mlops.workflow.fail_step import FailStep
from sagemaker.mlops.workflow.inference_component_step import InferenceComponentStep
from sagemaker.mlops.workflow.lambda_step import LambdaStep, LambdaOutput
from sagemaker.mlops.workflow.lineage_step import LineageStep
from sagemaker.mlops.workflow.model_step import ModelStep
from sagemaker.mlops.workflow.monitor_batch_transform_step import MonitorBatchTransformStep
from sagemaker.mlops.workflow.notebook_job_step import NotebookJobStep
Expand DownExpand Up@@ -98,9 +102,13 @@
"ConditionStep",
"EMRStep",
"EMRStepConfig",
"EndpointConfigStep",
"EndpointStep",
"FailStep",
"InferenceComponentStep",
"LambdaStep",
"LambdaOutput",
"LineageStep",
"ModelStep",
"MonitorBatchTransformStep",
"NotebookJobStep",
Expand Down
203 changes: 203 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/endpoint_step.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,203 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Step definitions for SageMaker Endpoint deployment in Pipelines.

These steps follow the ``step_args`` convention used by ``TrainingStep``
and ``ModelStep``: call the corresponding session method under a
:class:`~sagemaker.core.workflow.pipeline_context.PipelineSession` and
pass the returned step arguments to the step. The request is captured at
call time and the service call is deferred to pipeline execution.

Example::

pipeline_session = PipelineSession()

config_step_args = pipeline_session.endpoint_from_production_variants(
name="my-endpoint-config",
production_variants=[...],
)
config_step = EndpointConfigStep(name="CreateConfig", step_args=config_step_args)

endpoint_step_args = pipeline_session.create_endpoint(
endpoint_name="my-endpoint",
config_name="my-endpoint-config",
)
endpoint_step = EndpointStep(name="CreateEndpoint", step_args=endpoint_step_args)
"""

from __future__ import absolute_import

from typing import List, Optional, Union

from sagemaker.core.helper.pipeline_variable import RequestType
from sagemaker.core.workflow.pipeline_context import _JobStepArguments
from sagemaker.core.workflow.properties import Properties
from sagemaker.core.workflow.utilities import validate_step_args_input

from sagemaker.mlops.workflow.retry import RetryPolicy
from sagemaker.mlops.workflow.step_collections import StepCollection
from sagemaker.mlops.workflow.steps import (
CacheConfig,
ConfigurableRetryStep,
Step,
StepTypeEnum,
)


class EndpointConfigStep(ConfigurableRetryStep):
"""Creates a SageMaker EndpointConfig within a pipeline.

Wraps the SageMaker ``CreateEndpointConfig`` API. The ``step_args``
must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.endpoint_from_production_variants`
on a ``PipelineSession``.

``EndpointConfig`` is structurally cacheable (``cache_config``) and
retryable (``retry_policies``).
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
retry_policies: Optional[List[RetryPolicy]] = None,
):
"""Construct an ``EndpointConfigStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from
``pipeline_session.endpoint_from_production_variants()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
retry_policies (List[RetryPolicy]): Optional retry policies.
"""
super().__init__(
name=name,
step_type=StepTypeEnum.ENDPOINT_CONFIG,
display_name=display_name,
description=description,
depends_on=depends_on,
retry_policies=retry_policies,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"endpoint_from_production_variants"},
error_message=(
"The step_args of EndpointConfigStep must be obtained from "
"pipeline_session.endpoint_from_production_variants()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointConfigOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint_config``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointConfigOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict


class EndpointStep(Step):
"""Creates or updates a SageMaker Endpoint within a pipeline.

Wraps the SageMaker ``CreateEndpoint``/``UpdateEndpoint`` API -- the
pipeline chooses create-vs-update based on endpoint existence. The
``step_args`` must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.create_endpoint`
on a ``PipelineSession``.

``Endpoint`` is structurally cacheable but not retryable at the
pipeline level.
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
):
"""Construct an ``EndpointStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from ``pipeline_session.create_endpoint()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
"""
super().__init__(
name=name,
display_name=display_name,
description=description,
step_type=StepTypeEnum.ENDPOINT,
depends_on=depends_on,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"create_endpoint"},
error_message=(
"The step_args of EndpointStep must be obtained from "
"pipeline_session.create_endpoint()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict
Loading
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
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions sagemaker-core/src/sagemaker/core/apiutils/_base_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -219,10 +219,29 @@ def with_boto(self, boto_dict):
)
return self

# Lineage entity creation methods whose requests are captured as pipeline
# step arguments when invoked under a ``PipelineSession`` (instead of
# calling the service). Used by ``sagemaker.mlops.workflow.LineageStep``.
_PIPELINE_CAPTURABLE_METHODS = frozenset(
{"create_action", "create_artifact", "create_context", "add_association"}
)

def _invoke_api(self, boto_method, boto_method_members):
"""Invoke a SageMaker API."""
api_values = {k: v for k, v in vars(self).items() if k in boto_method_members}
api_kwargs = self.to_boto(api_values)

if boto_method in self._PIPELINE_CAPTURABLE_METHODS:
# Lazy import to avoid a circular dependency at module load time.
from sagemaker.core.workflow.pipeline_context import (
PipelineSession,
_JobStepArguments,
)

if isinstance(self.sagemaker_session, PipelineSession):
self.sagemaker_session.context = _JobStepArguments(boto_method, api_kwargs)
return self.sagemaker_session.context

api_method = getattr(self.sagemaker_session.sagemaker_client, boto_method)
api_boto_response = api_method(**api_kwargs)
return self.with_boto(api_boto_response)
73 changes: 57 additions & 16 deletions sagemaker-core/src/sagemaker/core/helper/session_helper.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1091,15 +1091,23 @@ def endpoint_from_production_variants(
config_options["ExecutionRoleArn"] = role

logger.info("Creating endpoint-config with name %s", name)
self.sagemaker_client.create_endpoint_config(**config_options)

return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,

def submit(request):
self.sagemaker_client.create_endpoint_config(**request)
return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,
)

result = self._intercept_create_request(
config_options, submit, self.endpoint_from_production_variants.__name__
)
if self._is_pipeline_context():
return self.context
return result

def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live_logging=False):
"""Create an Amazon SageMaker ``Endpoint`` according to the configuration in the request.
Expand DownExpand Up@@ -1129,16 +1137,28 @@ def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live
tags = self._append_sagemaker_config_tags(
tags, "{}.{}.{}".format(SAGEMAKER, ENDPOINT, TAGS)
)
try:
res = self.sagemaker_client.create_endpoint(
EndpointName=endpoint_name, EndpointConfigName=config_name, Tags=tags
)
create_endpoint_request = {
"EndpointName": endpoint_name,
"EndpointConfigName": config_name,
"Tags": tags,
}

def submit(request):
res = self.sagemaker_client.create_endpoint(**request)
if res:
self.endpoint_arn = res["EndpointArn"]

if wait:
self.wait_for_endpoint(endpoint_name, live_logging=live_logging)
return endpoint_name

try:
result = self._intercept_create_request(
create_endpoint_request, submit, self.create_endpoint.__name__
)
if self._is_pipeline_context():
return self.context
return result
except Exception as e:
troubleshooting = (
"https://docs.aws.amazon.com/sagemaker/latest/dg/"
Expand DownExpand Up@@ -1257,10 +1277,18 @@ def create_inference_component(
if tags and len(tags) != 0:
request["Tags"] = tags

self.sagemaker_client.create_inference_component(**request)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name
def submit(req):
self.sagemaker_client.create_inference_component(**req)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name

result = self._intercept_create_request(
request, submit, self.create_inference_component.__name__
)
if self._is_pipeline_context():
return self.context
return result

def wait_for_inference_component(self, inference_component_name, poll=20):
"""Wait for an Amazon SageMaker ``Inference Component`` deployment to complete.
Expand DownExpand Up@@ -1439,6 +1467,19 @@ def _intercept_create_request(
"""
return create(request)

def _is_pipeline_context(self) -> bool:
"""Whether this session is a pipeline session capturing requests.

Producer methods that support composing pipeline steps use this to
return the captured step arguments (``self.context``) instead of the
result of a service call. Always ``False`` for a plain ``Session``.
"""
# Lazy import to avoid a circular dependency: pipeline_context imports
# from this module at import time.
from sagemaker.core.workflow.pipeline_context import PipelineSession

return isinstance(self, PipelineSession)

def _create_inference_recommendations_job_request(
self,
role: str,
Expand Down
8 changes: 8 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/__init__.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,7 @@
functions, conditions, properties) and can import from sagemaker.train and sagemaker.serve
for orchestration purposes.
"""

from __future__ import absolute_import

__version__ = "0.1.0"
Expand DownExpand Up@@ -46,8 +47,11 @@
from sagemaker.mlops.workflow.clarify_check_step import ClarifyCheckStep
from sagemaker.mlops.workflow.condition_step import ConditionStep
from sagemaker.mlops.workflow.emr_step import EMRStep, EMRStepConfig
from sagemaker.mlops.workflow.endpoint_step import EndpointConfigStep, EndpointStep
from sagemaker.mlops.workflow.fail_step import FailStep
from sagemaker.mlops.workflow.inference_component_step import InferenceComponentStep
from sagemaker.mlops.workflow.lambda_step import LambdaStep, LambdaOutput
from sagemaker.mlops.workflow.lineage_step import LineageStep
from sagemaker.mlops.workflow.model_step import ModelStep
from sagemaker.mlops.workflow.monitor_batch_transform_step import MonitorBatchTransformStep
from sagemaker.mlops.workflow.notebook_job_step import NotebookJobStep
Expand DownExpand Up@@ -98,9 +102,13 @@
"ConditionStep",
"EMRStep",
"EMRStepConfig",
"EndpointConfigStep",
"EndpointStep",
"FailStep",
"InferenceComponentStep",
"LambdaStep",
"LambdaOutput",
"LineageStep",
"ModelStep",
"MonitorBatchTransformStep",
"NotebookJobStep",
Expand Down
203 changes: 203 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/endpoint_step.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,203 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Step definitions for SageMaker Endpoint deployment in Pipelines.

These steps follow the ``step_args`` convention used by ``TrainingStep``
and ``ModelStep``: call the corresponding session method under a
:class:`~sagemaker.core.workflow.pipeline_context.PipelineSession` and
pass the returned step arguments to the step. The request is captured at
call time and the service call is deferred to pipeline execution.

Example::

pipeline_session = PipelineSession()

config_step_args = pipeline_session.endpoint_from_production_variants(
name="my-endpoint-config",
production_variants=[...],
)
config_step = EndpointConfigStep(name="CreateConfig", step_args=config_step_args)

endpoint_step_args = pipeline_session.create_endpoint(
endpoint_name="my-endpoint",
config_name="my-endpoint-config",
)
endpoint_step = EndpointStep(name="CreateEndpoint", step_args=endpoint_step_args)
"""

from __future__ import absolute_import

from typing import List, Optional, Union

from sagemaker.core.helper.pipeline_variable import RequestType
from sagemaker.core.workflow.pipeline_context import _JobStepArguments
from sagemaker.core.workflow.properties import Properties
from sagemaker.core.workflow.utilities import validate_step_args_input

from sagemaker.mlops.workflow.retry import RetryPolicy
from sagemaker.mlops.workflow.step_collections import StepCollection
from sagemaker.mlops.workflow.steps import (
CacheConfig,
ConfigurableRetryStep,
Step,
StepTypeEnum,
)


class EndpointConfigStep(ConfigurableRetryStep):
"""Creates a SageMaker EndpointConfig within a pipeline.

Wraps the SageMaker ``CreateEndpointConfig`` API. The ``step_args``
must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.endpoint_from_production_variants`
on a ``PipelineSession``.

``EndpointConfig`` is structurally cacheable (``cache_config``) and
retryable (``retry_policies``).
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
retry_policies: Optional[List[RetryPolicy]] = None,
):
"""Construct an ``EndpointConfigStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from
``pipeline_session.endpoint_from_production_variants()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
retry_policies (List[RetryPolicy]): Optional retry policies.
"""
super().__init__(
name=name,
step_type=StepTypeEnum.ENDPOINT_CONFIG,
display_name=display_name,
description=description,
depends_on=depends_on,
retry_policies=retry_policies,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"endpoint_from_production_variants"},
error_message=(
"The step_args of EndpointConfigStep must be obtained from "
"pipeline_session.endpoint_from_production_variants()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointConfigOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint_config``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointConfigOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict


class EndpointStep(Step):
"""Creates or updates a SageMaker Endpoint within a pipeline.

Wraps the SageMaker ``CreateEndpoint``/``UpdateEndpoint`` API -- the
pipeline chooses create-vs-update based on endpoint existence. The
``step_args`` must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.create_endpoint`
on a ``PipelineSession``.

``Endpoint`` is structurally cacheable but not retryable at the
pipeline level.
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
):
"""Construct an ``EndpointStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from ``pipeline_session.create_endpoint()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
"""
super().__init__(
name=name,
display_name=display_name,
description=description,
step_type=StepTypeEnum.ENDPOINT,
depends_on=depends_on,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"create_endpoint"},
error_message=(
"The step_args of EndpointStep must be obtained from "
"pipeline_session.create_endpoint()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict
Loading
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
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions sagemaker-core/src/sagemaker/core/apiutils/_base_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -219,10 +219,29 @@ def with_boto(self, boto_dict):
)
return self

# Lineage entity creation methods whose requests are captured as pipeline
# step arguments when invoked under a ``PipelineSession`` (instead of
# calling the service). Used by ``sagemaker.mlops.workflow.LineageStep``.
_PIPELINE_CAPTURABLE_METHODS = frozenset(
{"create_action", "create_artifact", "create_context", "add_association"}
)

def _invoke_api(self, boto_method, boto_method_members):
"""Invoke a SageMaker API."""
api_values = {k: v for k, v in vars(self).items() if k in boto_method_members}
api_kwargs = self.to_boto(api_values)

if boto_method in self._PIPELINE_CAPTURABLE_METHODS:
# Lazy import to avoid a circular dependency at module load time.
from sagemaker.core.workflow.pipeline_context import (
PipelineSession,
_JobStepArguments,
)

if isinstance(self.sagemaker_session, PipelineSession):
self.sagemaker_session.context = _JobStepArguments(boto_method, api_kwargs)
return self.sagemaker_session.context

api_method = getattr(self.sagemaker_session.sagemaker_client, boto_method)
api_boto_response = api_method(**api_kwargs)
return self.with_boto(api_boto_response)
73 changes: 57 additions & 16 deletions sagemaker-core/src/sagemaker/core/helper/session_helper.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1091,15 +1091,23 @@ def endpoint_from_production_variants(
config_options["ExecutionRoleArn"] = role

logger.info("Creating endpoint-config with name %s", name)
self.sagemaker_client.create_endpoint_config(**config_options)

return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,

def submit(request):
self.sagemaker_client.create_endpoint_config(**request)
return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,
)

result = self._intercept_create_request(
config_options, submit, self.endpoint_from_production_variants.__name__
)
if self._is_pipeline_context():
return self.context
return result

def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live_logging=False):
"""Create an Amazon SageMaker ``Endpoint`` according to the configuration in the request.
Expand DownExpand Up@@ -1129,16 +1137,28 @@ def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live
tags = self._append_sagemaker_config_tags(
tags, "{}.{}.{}".format(SAGEMAKER, ENDPOINT, TAGS)
)
try:
res = self.sagemaker_client.create_endpoint(
EndpointName=endpoint_name, EndpointConfigName=config_name, Tags=tags
)
create_endpoint_request = {
"EndpointName": endpoint_name,
"EndpointConfigName": config_name,
"Tags": tags,
}

def submit(request):
res = self.sagemaker_client.create_endpoint(**request)
if res:
self.endpoint_arn = res["EndpointArn"]

if wait:
self.wait_for_endpoint(endpoint_name, live_logging=live_logging)
return endpoint_name

try:
result = self._intercept_create_request(
create_endpoint_request, submit, self.create_endpoint.__name__
)
if self._is_pipeline_context():
return self.context
return result
except Exception as e:
troubleshooting = (
"https://docs.aws.amazon.com/sagemaker/latest/dg/"
Expand DownExpand Up@@ -1257,10 +1277,18 @@ def create_inference_component(
if tags and len(tags) != 0:
request["Tags"] = tags

self.sagemaker_client.create_inference_component(**request)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name
def submit(req):
self.sagemaker_client.create_inference_component(**req)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name

result = self._intercept_create_request(
request, submit, self.create_inference_component.__name__
)
if self._is_pipeline_context():
return self.context
return result

def wait_for_inference_component(self, inference_component_name, poll=20):
"""Wait for an Amazon SageMaker ``Inference Component`` deployment to complete.
Expand DownExpand Up@@ -1439,6 +1467,19 @@ def _intercept_create_request(
"""
return create(request)

def _is_pipeline_context(self) -> bool:
"""Whether this session is a pipeline session capturing requests.

Producer methods that support composing pipeline steps use this to
return the captured step arguments (``self.context``) instead of the
result of a service call. Always ``False`` for a plain ``Session``.
"""
# Lazy import to avoid a circular dependency: pipeline_context imports
# from this module at import time.
from sagemaker.core.workflow.pipeline_context import PipelineSession

return isinstance(self, PipelineSession)

def _create_inference_recommendations_job_request(
self,
role: str,
Expand Down
8 changes: 8 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/__init__.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,7 @@
functions, conditions, properties) and can import from sagemaker.train and sagemaker.serve
for orchestration purposes.
"""

from __future__ import absolute_import

__version__ = "0.1.0"
Expand DownExpand Up@@ -46,8 +47,11 @@
from sagemaker.mlops.workflow.clarify_check_step import ClarifyCheckStep
from sagemaker.mlops.workflow.condition_step import ConditionStep
from sagemaker.mlops.workflow.emr_step import EMRStep, EMRStepConfig
from sagemaker.mlops.workflow.endpoint_step import EndpointConfigStep, EndpointStep
from sagemaker.mlops.workflow.fail_step import FailStep
from sagemaker.mlops.workflow.inference_component_step import InferenceComponentStep
from sagemaker.mlops.workflow.lambda_step import LambdaStep, LambdaOutput
from sagemaker.mlops.workflow.lineage_step import LineageStep
from sagemaker.mlops.workflow.model_step import ModelStep
from sagemaker.mlops.workflow.monitor_batch_transform_step import MonitorBatchTransformStep
from sagemaker.mlops.workflow.notebook_job_step import NotebookJobStep
Expand DownExpand Up@@ -98,9 +102,13 @@
"ConditionStep",
"EMRStep",
"EMRStepConfig",
"EndpointConfigStep",
"EndpointStep",
"FailStep",
"InferenceComponentStep",
"LambdaStep",
"LambdaOutput",
"LineageStep",
"ModelStep",
"MonitorBatchTransformStep",
"NotebookJobStep",
Expand Down
203 changes: 203 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/endpoint_step.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,203 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Step definitions for SageMaker Endpoint deployment in Pipelines.

These steps follow the ``step_args`` convention used by ``TrainingStep``
and ``ModelStep``: call the corresponding session method under a
:class:`~sagemaker.core.workflow.pipeline_context.PipelineSession` and
pass the returned step arguments to the step. The request is captured at
call time and the service call is deferred to pipeline execution.

Example::

pipeline_session = PipelineSession()

config_step_args = pipeline_session.endpoint_from_production_variants(
name="my-endpoint-config",
production_variants=[...],
)
config_step = EndpointConfigStep(name="CreateConfig", step_args=config_step_args)

endpoint_step_args = pipeline_session.create_endpoint(
endpoint_name="my-endpoint",
config_name="my-endpoint-config",
)
endpoint_step = EndpointStep(name="CreateEndpoint", step_args=endpoint_step_args)
"""

from __future__ import absolute_import

from typing import List, Optional, Union

from sagemaker.core.helper.pipeline_variable import RequestType
from sagemaker.core.workflow.pipeline_context import _JobStepArguments
from sagemaker.core.workflow.properties import Properties
from sagemaker.core.workflow.utilities import validate_step_args_input

from sagemaker.mlops.workflow.retry import RetryPolicy
from sagemaker.mlops.workflow.step_collections import StepCollection
from sagemaker.mlops.workflow.steps import (
CacheConfig,
ConfigurableRetryStep,
Step,
StepTypeEnum,
)


class EndpointConfigStep(ConfigurableRetryStep):
"""Creates a SageMaker EndpointConfig within a pipeline.

Wraps the SageMaker ``CreateEndpointConfig`` API. The ``step_args``
must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.endpoint_from_production_variants`
on a ``PipelineSession``.

``EndpointConfig`` is structurally cacheable (``cache_config``) and
retryable (``retry_policies``).
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
retry_policies: Optional[List[RetryPolicy]] = None,
):
"""Construct an ``EndpointConfigStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from
``pipeline_session.endpoint_from_production_variants()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
retry_policies (List[RetryPolicy]): Optional retry policies.
"""
super().__init__(
name=name,
step_type=StepTypeEnum.ENDPOINT_CONFIG,
display_name=display_name,
description=description,
depends_on=depends_on,
retry_policies=retry_policies,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"endpoint_from_production_variants"},
error_message=(
"The step_args of EndpointConfigStep must be obtained from "
"pipeline_session.endpoint_from_production_variants()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointConfigOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint_config``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointConfigOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict


class EndpointStep(Step):
"""Creates or updates a SageMaker Endpoint within a pipeline.

Wraps the SageMaker ``CreateEndpoint``/``UpdateEndpoint`` API -- the
pipeline chooses create-vs-update based on endpoint existence. The
``step_args`` must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.create_endpoint`
on a ``PipelineSession``.

``Endpoint`` is structurally cacheable but not retryable at the
pipeline level.
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
):
"""Construct an ``EndpointStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from ``pipeline_session.create_endpoint()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
"""
super().__init__(
name=name,
display_name=display_name,
description=description,
step_type=StepTypeEnum.ENDPOINT,
depends_on=depends_on,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"create_endpoint"},
error_message=(
"The step_args of EndpointStep must be obtained from "
"pipeline_session.create_endpoint()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict
Loading
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
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 19 additions & 0 deletions sagemaker-core/src/sagemaker/core/apiutils/_base_types.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -219,10 +219,29 @@ def with_boto(self, boto_dict):
)
return self

# Lineage entity creation methods whose requests are captured as pipeline
# step arguments when invoked under a ``PipelineSession`` (instead of
# calling the service). Used by ``sagemaker.mlops.workflow.LineageStep``.
_PIPELINE_CAPTURABLE_METHODS = frozenset(
{"create_action", "create_artifact", "create_context", "add_association"}
)

def _invoke_api(self, boto_method, boto_method_members):
"""Invoke a SageMaker API."""
api_values = {k: v for k, v in vars(self).items() if k in boto_method_members}
api_kwargs = self.to_boto(api_values)

if boto_method in self._PIPELINE_CAPTURABLE_METHODS:
# Lazy import to avoid a circular dependency at module load time.
from sagemaker.core.workflow.pipeline_context import (
PipelineSession,
_JobStepArguments,
)

if isinstance(self.sagemaker_session, PipelineSession):
self.sagemaker_session.context = _JobStepArguments(boto_method, api_kwargs)
return self.sagemaker_session.context

api_method = getattr(self.sagemaker_session.sagemaker_client, boto_method)
api_boto_response = api_method(**api_kwargs)
return self.with_boto(api_boto_response)
73 changes: 57 additions & 16 deletions sagemaker-core/src/sagemaker/core/helper/session_helper.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1091,15 +1091,23 @@ def endpoint_from_production_variants(
config_options["ExecutionRoleArn"] = role

logger.info("Creating endpoint-config with name %s", name)
self.sagemaker_client.create_endpoint_config(**config_options)

return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,

def submit(request):
self.sagemaker_client.create_endpoint_config(**request)
return self.create_endpoint(
endpoint_name=name,
config_name=name,
tags=endpoint_tags,
wait=wait,
live_logging=live_logging,
)

result = self._intercept_create_request(
config_options, submit, self.endpoint_from_production_variants.__name__
)
if self._is_pipeline_context():
return self.context
return result

def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live_logging=False):
"""Create an Amazon SageMaker ``Endpoint`` according to the configuration in the request.
Expand DownExpand Up@@ -1129,16 +1137,28 @@ def create_endpoint(self, endpoint_name, config_name, tags=None, wait=True, live
tags = self._append_sagemaker_config_tags(
tags, "{}.{}.{}".format(SAGEMAKER, ENDPOINT, TAGS)
)
try:
res = self.sagemaker_client.create_endpoint(
EndpointName=endpoint_name, EndpointConfigName=config_name, Tags=tags
)
create_endpoint_request = {
"EndpointName": endpoint_name,
"EndpointConfigName": config_name,
"Tags": tags,
}

def submit(request):
res = self.sagemaker_client.create_endpoint(**request)
if res:
self.endpoint_arn = res["EndpointArn"]

if wait:
self.wait_for_endpoint(endpoint_name, live_logging=live_logging)
return endpoint_name

try:
result = self._intercept_create_request(
create_endpoint_request, submit, self.create_endpoint.__name__
)
if self._is_pipeline_context():
return self.context
return result
except Exception as e:
troubleshooting = (
"https://docs.aws.amazon.com/sagemaker/latest/dg/"
Expand DownExpand Up@@ -1257,10 +1277,18 @@ def create_inference_component(
if tags and len(tags) != 0:
request["Tags"] = tags

self.sagemaker_client.create_inference_component(**request)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name
def submit(req):
self.sagemaker_client.create_inference_component(**req)
if wait:
self.wait_for_inference_component(inference_component_name)
return inference_component_name

result = self._intercept_create_request(
request, submit, self.create_inference_component.__name__
)
if self._is_pipeline_context():
return self.context
return result

def wait_for_inference_component(self, inference_component_name, poll=20):
"""Wait for an Amazon SageMaker ``Inference Component`` deployment to complete.
Expand DownExpand Up@@ -1439,6 +1467,19 @@ def _intercept_create_request(
"""
return create(request)

def _is_pipeline_context(self) -> bool:
"""Whether this session is a pipeline session capturing requests.

Producer methods that support composing pipeline steps use this to
return the captured step arguments (``self.context``) instead of the
result of a service call. Always ``False`` for a plain ``Session``.
"""
# Lazy import to avoid a circular dependency: pipeline_context imports
# from this module at import time.
from sagemaker.core.workflow.pipeline_context import PipelineSession

return isinstance(self, PipelineSession)

def _create_inference_recommendations_job_request(
self,
role: str,
Expand Down
8 changes: 8 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/__init__.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -14,6 +14,7 @@
functions, conditions, properties) and can import from sagemaker.train and sagemaker.serve
for orchestration purposes.
"""

from __future__ import absolute_import

__version__ = "0.1.0"
Expand DownExpand Up@@ -46,8 +47,11 @@
from sagemaker.mlops.workflow.clarify_check_step import ClarifyCheckStep
from sagemaker.mlops.workflow.condition_step import ConditionStep
from sagemaker.mlops.workflow.emr_step import EMRStep, EMRStepConfig
from sagemaker.mlops.workflow.endpoint_step import EndpointConfigStep, EndpointStep
from sagemaker.mlops.workflow.fail_step import FailStep
from sagemaker.mlops.workflow.inference_component_step import InferenceComponentStep
from sagemaker.mlops.workflow.lambda_step import LambdaStep, LambdaOutput
from sagemaker.mlops.workflow.lineage_step import LineageStep
from sagemaker.mlops.workflow.model_step import ModelStep
from sagemaker.mlops.workflow.monitor_batch_transform_step import MonitorBatchTransformStep
from sagemaker.mlops.workflow.notebook_job_step import NotebookJobStep
Expand DownExpand Up@@ -98,9 +102,13 @@
"ConditionStep",
"EMRStep",
"EMRStepConfig",
"EndpointConfigStep",
"EndpointStep",
"FailStep",
"InferenceComponentStep",
"LambdaStep",
"LambdaOutput",
"LineageStep",
"ModelStep",
"MonitorBatchTransformStep",
"NotebookJobStep",
Expand Down
203 changes: 203 additions & 0 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/endpoint_step.py
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,203 @@
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License"). You
# may not use this file except in compliance with the License. A copy of
# the License is located at
#
# http://aws.amazon.com/apache2.0/
#
# or in the "license" file accompanying this file. This file is
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
# ANY KIND, either express or implied. See the License for the specific
# language governing permissions and limitations under the License.
"""Step definitions for SageMaker Endpoint deployment in Pipelines.

These steps follow the ``step_args`` convention used by ``TrainingStep``
and ``ModelStep``: call the corresponding session method under a
:class:`~sagemaker.core.workflow.pipeline_context.PipelineSession` and
pass the returned step arguments to the step. The request is captured at
call time and the service call is deferred to pipeline execution.

Example::

pipeline_session = PipelineSession()

config_step_args = pipeline_session.endpoint_from_production_variants(
name="my-endpoint-config",
production_variants=[...],
)
config_step = EndpointConfigStep(name="CreateConfig", step_args=config_step_args)

endpoint_step_args = pipeline_session.create_endpoint(
endpoint_name="my-endpoint",
config_name="my-endpoint-config",
)
endpoint_step = EndpointStep(name="CreateEndpoint", step_args=endpoint_step_args)
"""

from __future__ import absolute_import

from typing import List, Optional, Union

from sagemaker.core.helper.pipeline_variable import RequestType
from sagemaker.core.workflow.pipeline_context import _JobStepArguments
from sagemaker.core.workflow.properties import Properties
from sagemaker.core.workflow.utilities import validate_step_args_input

from sagemaker.mlops.workflow.retry import RetryPolicy
from sagemaker.mlops.workflow.step_collections import StepCollection
from sagemaker.mlops.workflow.steps import (
CacheConfig,
ConfigurableRetryStep,
Step,
StepTypeEnum,
)


class EndpointConfigStep(ConfigurableRetryStep):
"""Creates a SageMaker EndpointConfig within a pipeline.

Wraps the SageMaker ``CreateEndpointConfig`` API. The ``step_args``
must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.endpoint_from_production_variants`
on a ``PipelineSession``.

``EndpointConfig`` is structurally cacheable (``cache_config``) and
retryable (``retry_policies``).
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
retry_policies: Optional[List[RetryPolicy]] = None,
):
"""Construct an ``EndpointConfigStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from
``pipeline_session.endpoint_from_production_variants()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
retry_policies (List[RetryPolicy]): Optional retry policies.
"""
super().__init__(
name=name,
step_type=StepTypeEnum.ENDPOINT_CONFIG,
display_name=display_name,
description=description,
depends_on=depends_on,
retry_policies=retry_policies,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"endpoint_from_production_variants"},
error_message=(
"The step_args of EndpointConfigStep must be obtained from "
"pipeline_session.endpoint_from_production_variants()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointConfigOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint_config``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointConfigOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict


class EndpointStep(Step):
"""Creates or updates a SageMaker Endpoint within a pipeline.

Wraps the SageMaker ``CreateEndpoint``/``UpdateEndpoint`` API -- the
pipeline chooses create-vs-update based on endpoint existence. The
``step_args`` must be obtained by calling
:meth:`~sagemaker.core.helper.session_helper.Session.create_endpoint`
on a ``PipelineSession``.

``Endpoint`` is structurally cacheable but not retryable at the
pipeline level.
"""

def __init__(
self,
name: str,
step_args: _JobStepArguments,
display_name: Optional[str] = None,
description: Optional[str] = None,
depends_on: Optional[List[Union[str, Step, StepCollection]]] = None,
cache_config: Optional[CacheConfig] = None,
):
"""Construct an ``EndpointStep``.

Args:
name (str): The name of the step.
step_args (_JobStepArguments): The arguments for this step,
obtained from ``pipeline_session.create_endpoint()``.
display_name (str): Optional display name.
description (str): Optional description.
depends_on (List[Union[str, Step, StepCollection]]): Optional
explicit step dependencies.
cache_config (CacheConfig): Optional cache configuration.
"""
super().__init__(
name=name,
display_name=display_name,
description=description,
step_type=StepTypeEnum.ENDPOINT,
depends_on=depends_on,
)
validate_step_args_input(
step_args=step_args,
expected_caller={"create_endpoint"},
error_message=(
"The step_args of EndpointStep must be obtained from "
"pipeline_session.create_endpoint()."
),
)
self.step_args = step_args
self.cache_config = cache_config
self._properties = Properties(
step_name=name, step=self, shape_name="DescribeEndpointOutput"
)

@property
def arguments(self) -> RequestType:
"""The arguments dictionary that is used to call ``create_endpoint``."""
return self.step_args.args

@property
def properties(self):
"""A ``Properties`` object shaped like ``DescribeEndpointOutput``."""
return self._properties

def to_request(self) -> RequestType:
"""Get the request structure for workflow service calls."""
request_dict = super().to_request()
if self.cache_config:
request_dict.update(self.cache_config.config)
return request_dict
Loading
Loading