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
98 changes: 74 additions & 24 deletions sagemaker-mlops/src/sagemaker/mlops/workflow/model_step.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
from typing import Union, List, Dict, Optional

from sagemaker.core.resources import Model
from sagemaker.core.shapes import OutputDataConfig
from sagemaker.mlops.workflow._utils import _RepackModelStep
from sagemaker.core.workflow.pipeline_context import PipelineSession, _ModelStepArguments
from sagemaker.mlops.workflow.retry import RetryPolicy, SageMakerJobStepRetryPolicy
Expand Down Expand Up @@ -211,9 +212,53 @@ 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,
"model_kms_key": getattr(model, "model_kms_key", None),
}
# 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),
"model_kms_key": getattr(model, "model_kms_key", 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

# ModelBuilder is what ModelBuilder.register()/.build() put into the pipeline
# context. The core ``Model`` arm is legacy/duck-typed: no v3 @runnable_by_pipeline
# path produces one, and sagemaker.core.resources.Model has no sagemaker_session /
# role / entry_point (GH #5829), so it is kept only for objects that happen to
# expose those attributes rather than as a supported v3 entry point.
if isinstance(self._model, (Model, ModelBuilder)):
model_list = [self._model]
else:
logger.warning("No models to repack")
Expand All @@ -224,27 +269,39 @@ 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
# Send the repacked artifact to the location ModelBuilder computed from the
# user's bucket / code_location, encrypted with their model KMS key. These
# reach ModelTrainer via _RepackModelStep's **kwargs. Without this the
# repacked tarball silently lands in ModelTrainer's default bucket with no
# CMK, which breaks accounts with a mandated bucket or SSE-KMS policy.
# setdefault so an explicit repack_model_step_settings override wins.
if self._runtime_repack_output_prefix or fields["model_kms_key"]:
self._repack_model_step_settings.setdefault(
"output_data_config",
OutputDataConfig(
s3_output_path=self._runtime_repack_output_prefix,
kms_key_id=fields["model_kms_key"],
),
)
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 +310,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 +345,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
151 changes: 150 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,152 @@ 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_repack_step_gets_output_location_and_kms_key():
"""The repacked artifact must go to ModelBuilder's runtime_repack_output_prefix and be
encrypted with the model's KMS key. Both reach ModelTrainer via _RepackModelStep's
**kwargs; dropping them sends the artifact to ModelTrainer's default bucket with no CMK,
which breaks accounts with a mandated bucket or an SSE-KMS bucket policy."""
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
builder.model_kms_key = "arn:aws:kms:us-west-2:111122223333:key/abc"

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:
ms.ModelStep(name="step", step_args=step_args)

_, kwargs = mock_repack.call_args
odc = kwargs["output_data_config"]
assert odc.s3_output_path == "s3://bucket/prefix"
assert odc.kms_key_id == "arn:aws:kms:us-west-2:111122223333:key/abc"


def test_repack_step_output_config_respects_user_override():
"""An explicit output_data_config in repack_model_step_settings must win."""
from sagemaker.serve.model_builder import ModelBuilder
from sagemaker.core.shapes import OutputDataConfig
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
builder.model_kms_key = "arn:aws:kms:us-west-2:111122223333:key/abc"

step_args = _FakeModelStepArgs(builder, {id(builder)})
fake_repack = Mock()
fake_repack.properties.ModelArtifacts.S3ModelArtifacts = "s3://repacked/model.tar.gz"
mine = OutputDataConfig(s3_output_path="s3://mine/out", kms_key_id="my-key")

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

_, kwargs = mock_repack.call_args
assert kwargs["output_data_config"] is mine


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