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
76 changes: 52 additions & 24 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/model_step.py
Original file line number Diff line number Diff line change
Expand Up @@ -211,9 +211,46 @@ def properties(self):
"""A Properties object representing the appropriate SageMaker response data model."""
return self._properties

@staticmethod
def _repack_inputs_for(model):
"""Return the repack inputs for ``model`` mapped to _RepackModelStep's params.
Handles both a v3 ``ModelBuilder`` (what ``ModelBuilder.register()``/``.build()``
place into the pipeline context) and a legacy ``sagemaker.core.resources.Model``.
A ``ModelBuilder`` exposes these under different names (``role_arn``,
``s3_model_data_url``, ``model_name``) and carries its inference requirements on
its ``source_code``; normalize them here. See GH #5828 / #5829.
"""
from sagemaker.serve.model_builder import ModelBuilder

if isinstance(model, ModelBuilder):
source_code = getattr(model, "source_code", None)
requirements = getattr(source_code, "requirements", None) if source_code else None
return {
"name": getattr(model, "model_name", None),
"sagemaker_session": model.sagemaker_session,
"role": getattr(model, "role_arn", None),
"model_data": getattr(model, "s3_model_data_url", None),
"entry_point": getattr(model, "entry_point", None),
"source_dir": getattr(model, "source_dir", None),
"requirements": requirements,
}
# Legacy sagemaker.core.resources.Model path.
return {
"name": getattr(model, "name", None),
"sagemaker_session": getattr(model, "sagemaker_session", None),
"role": getattr(model, "role", None),
"model_data": getattr(model, "model_data", None),
"entry_point": getattr(model, "entry_point", None),
"source_dir": getattr(model, "source_dir", None),
"requirements": getattr(model, "requirements", None),
}

def _append_repack_model_step(self):
"""Create and append a `_RepackModelStep` for the runtime repack"""
if isinstance(self._model, Model):
from sagemaker.serve.model_builder import ModelBuilder

if isinstance(self._model, (Model, ModelBuilder)):
model_list = [self._model]
else:
logger.warning("No models to repack")
Expand All @@ -224,27 +261,25 @@ def _append_repack_model_step(self):
security_group_ids, subnets = self._resolve_repack_model_step_vpc_configs()

for i, model in enumerate(model_list):
# need_runtime_repack holds the id() of the original model/builder object,
# so the membership test must run against ``model`` itself, not a wrapper.
runtime_repack_flg = (
self._need_runtime_repack and id(model) in self._need_runtime_repack
)
if runtime_repack_flg:
name_base = model.name or i
fields = self._repack_inputs_for(model)
name_base = fields["name"] or i
repack_model_step = _RepackModelStep(
name="{}-{}-{}".format(self.name, _REPACK_MODEL_NAME_BASE, name_base),
sagemaker_session=(
self._repack_model_step_settings.pop("sagemaker_session", None)
or self._model.sagemaker_session
or model.sagemaker_session
or fields["sagemaker_session"]
),
role=(
self._repack_model_step_settings.pop("role", None)
or self._model.role
or model.role
),
model_data=model.model_data,
entry_point=model.entry_point,
source_dir=model.source_dir,
dependencies=model.dependencies,
role=(self._repack_model_step_settings.pop("role", None) or fields["role"]),
model_data=fields["model_data"],
entry_point=fields["entry_point"],
source_dir=fields["source_dir"],
requirements=fields["requirements"],
subnets=subnets,
security_group_ids=security_group_ids,
description=(
Expand All @@ -253,14 +288,6 @@ def _append_repack_model_step(self):
),
depends_on=self.depends_on,
retry_policies=self._repack_model_retry_policies,
output_path=(
self._repack_model_step_settings.pop("output_path", None)
or self._runtime_repack_output_prefix
),
output_kms_key=(
self._repack_model_step_settings.pop("output_kms_key", None)
or model.model_kms_key
),
**self._repack_model_step_settings,
)
self.steps.append(repack_model_step)
Expand Down Expand Up @@ -296,9 +323,10 @@ def _resolve_repack_model_step_vpc_configs(self):
subnets = self._repack_model_step_settings.pop("subnets", None)
return security_group_ids, subnets

if self._model.vpc_config:
security_group_ids = self._model.vpc_config.get("SecurityGroupIds", None)
subnets = self._model.vpc_config.get("Subnets", None)
vpc_config = getattr(self._model, "vpc_config", None)
if vpc_config:
security_group_ids = vpc_config.get("SecurityGroupIds", None)
subnets = vpc_config.get("Subnets", None)
return security_group_ids, subnets

return None, None
82 changes: 81 additions & 1 deletion sagemaker-mlops/tests/unit/workflow/test_model_step.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@

from __future__ import absolute_import

from unittest.mock import patch
from unittest.mock import Mock, patch


def test_model_step_properties():
Expand All @@ -27,3 +27,83 @@ def test_model_step_properties():
step = ModelStep(name="model-step", step_args=step_args)
assert step.name == "model-step"
assert hasattr(step, "properties")


def _pipeline_session():
from sagemaker.core.workflow.pipeline_context import PipelineSession

ps = Mock(spec=PipelineSession)
ps.context = Mock()
ps.boto_region_name = "us-west-2"
return ps


class _FakeModelStepArgs:
"""Mimics the _ModelStepArguments produced by ModelBuilder.register() under a
PipelineSession (a register/create_model_package request that needs a repack)."""

def __init__(self, model, need_runtime_repack):
self.model = model
self.need_runtime_repack = need_runtime_repack
self.runtime_repack_output_prefix = "s3://bucket/prefix"
self.create_model_request = None
self.create_model_package_request = {
"InferenceSpecification": {"Containers": [{"ModelDataUrl": "s3://orig/model.tar.gz"}]}
}


def test_model_builder_register_appends_repack_step():
"""GH #5828/#5829: ModelBuilder.register() in a ModelStep must emit a repack step
(v2 parity) and rewire the container ModelDataUrl to the repacked artifact."""
from sagemaker.serve.model_builder import ModelBuilder
from sagemaker.mlops.workflow import model_step as ms

ps = _pipeline_session()
builder = Mock(spec=ModelBuilder)
builder.sagemaker_session = ps
builder.model_name = "my-model"
builder.role_arn = "arn:aws:iam::111122223333:role/R"
builder.s3_model_data_url = "s3://orig/model.tar.gz"
builder.entry_point = "inference.py"
builder.source_dir = "/code"
builder.source_code = Mock(requirements="requirements.txt")
builder.vpc_config = None

step_args = _FakeModelStepArgs(builder, {id(builder)})
fake_repack = Mock()
fake_repack.properties.ModelArtifacts.S3ModelArtifacts = "s3://repacked/model.tar.gz"

with patch("sagemaker.core.workflow.utilities.validate_step_args_input"):
with patch.object(ms, "_RepackModelStep", return_value=fake_repack) as mock_repack:
step = ms.ModelStep(name="step", step_args=step_args)

# A repack step was generated (the bug: none was on master).
assert len(step.steps) == 1
# ModelBuilder attributes were mapped to the repack step's parameters.
_, kwargs = mock_repack.call_args
assert kwargs["role"] == "arn:aws:iam::111122223333:role/R"
assert kwargs["model_data"] == "s3://orig/model.tar.gz"
assert kwargs["entry_point"] == "inference.py"
assert kwargs["source_dir"] == "/code"
assert kwargs["requirements"] == "requirements.txt"
assert kwargs["sagemaker_session"] is ps
# The container now points at the repacked artifact.
container = step_args.create_model_package_request["InferenceSpecification"]["Containers"][0]
assert container["ModelDataUrl"] == "s3://repacked/model.tar.gz"


def test_no_repack_step_for_unrecognized_model_type():
"""An object that is neither a core Model nor a ModelBuilder yields no repack step."""
from sagemaker.mlops.workflow import model_step as ms

ps = _pipeline_session()
unknown = Mock()
unknown.sagemaker_session = ps
step_args = _FakeModelStepArgs(unknown, {id(unknown)})

with patch("sagemaker.core.workflow.utilities.validate_step_args_input"):
with patch.object(ms, "_RepackModelStep") as mock_repack:
step = ms.ModelStep(name="step", step_args=step_args)

assert step.steps == []
mock_repack.assert_not_called()
Loading