Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 40 additions & 2 deletions sagemaker-serve/src/sagemaker/serve/model_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -1879,6 +1879,42 @@ def _build_for_passthrough(self) -> Model:

self.secret_key = ""

model_artifact_uri = None
if self.model_path and self.model_path.startswith("s3://"):
model_artifact_uri = self.model_path
elif isinstance(self.s3_model_data_url, str) and self.s3_model_data_url.startswith(
"s3://"
):
model_artifact_uri = self.s3_model_data_url

has_source_code = bool(
getattr(self, "entry_point", None) and getattr(self, "source_dir", None)
)

# Repack source_code into the model artifact so build() produces a
# self-contained model.tar.gz (code under code/).
if has_source_code and model_artifact_uri:
if not (
isinstance(self.s3_model_data_url, str)
and self.s3_model_data_url.startswith("s3://")
):
self.s3_model_data_url = model_artifact_uri
self.s3_upload_path = None

if self.mode in LOCAL_MODES:
self._prepare_for_mode()

return self._create_model()

if getattr(self, "entry_point", None) and model_artifact_uri:
# entry_point provided without a source_dir: repack cannot bundle the
# code, so it would be dropped. Warn instead of silently ignoring it.
logger.warning(
"source_code was provided without a source_dir; the inference code "
"will not be repacked into the model artifact. Provide "
"SourceCode(source_dir=...) to bundle custom inference code."
)

if self.model_path and self.model_path.startswith("s3://"):
self.s3_upload_path = self.model_path
else:
Expand Down Expand Up @@ -2282,21 +2318,23 @@ def _upload_code(self, key_prefix: str, repack: bool = False) -> None:
script_name=os.path.basename(self.entry_point),
)

repack_dependencies = self.script_dependencies or []

logger.info(
"Repacking model artifact (%s), script artifact "
"(%s), and dependencies (%s) "
"into single tar.gz file located at %s. "
"This may take some time depending on model size...",
self.s3_model_data_url,
self.source_dir,
self.dependencies,
repack_dependencies,
repacked_model_data,
)

repack_model(
inference_script=self.entry_point,
source_directory=self.source_dir,
dependencies=self.dependencies,
dependencies=repack_dependencies,
model_uri=self.s3_model_data_url,
repacked_model_uri=repacked_model_data,
sagemaker_session=self.sagemaker_session,
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,142 @@
# 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.
"""Integration test for ModelBuilder passthrough source_code repack.

Covers the fix where an image_uri build with a model artifact and custom
source_code repacks the code into the artifact (instead of silently dropping
it). This calls build() only (no deploy) so it runs in seconds.
"""
from __future__ import absolute_import

import io
import os
import tarfile
import tempfile
import uuid

import boto3
import pytest

from sagemaker.serve.model_builder import ModelBuilder
from sagemaker.serve.mode.function_pointers import Mode
from sagemaker.core.training.configs import SourceCode
from sagemaker.core.helper.session_helper import Session, get_execution_role

MODEL_NAME_PREFIX = "mb-passthrough-repack"


def _upload_raw_artifact(s3_client, bucket, prefix):
"""Upload a minimal model.tar.gz containing only a model file (no code/)."""
tmp = tempfile.mkdtemp()
with open(os.path.join(tmp, "model.json"), "w") as f:
f.write('{"weights": [1, 1, 1, 1]}')
tar_path = os.path.join(tmp, "model.tar.gz")
with tarfile.open(tar_path, "w:gz") as t:
t.add(os.path.join(tmp, "model.json"), arcname="model.json")
key = f"{prefix}/model.tar.gz"
s3_client.upload_file(tar_path, bucket, key)
return f"s3://{bucket}/{key}", key


def _make_source_code_dir():
d = tempfile.mkdtemp()
with open(os.path.join(d, "inference.py"), "w") as f:
f.write(
"def model_fn(model_dir):\n return None\n"
"def predict_fn(data, model):\n return data\n"
)
with open(os.path.join(d, "requirements.txt"), "w") as f:
f.write("joblib\n")
return d


def _tar_members(s3_client, s3_uri):
_, _, rest = s3_uri.partition("s3://")
bucket, _, key = rest.partition("/")
body = s3_client.get_object(Bucket=bucket, Key=key)["Body"].read()
return tarfile.open(fileobj=io.BytesIO(body), mode="r:gz").getnames()


@pytest.mark.slow_test
def test_build_repacks_source_code_into_artifact():
"""build() with image_uri + model artifact + source_code repacks code/ into
the model.tar.gz. No deploy - runs in seconds."""
session = Session()
region = session.boto_region_name
bucket = session.default_bucket()
role = get_execution_role(sagemaker_session=session)
s3_client = boto3.client("s3", region_name=region)

unique_id = uuid.uuid4().hex[:8]
prefix = f"{MODEL_NAME_PREFIX}/{unique_id}"
s3_keys = []
core_model = None

try:
from sagemaker.core import image_uris

image_uri = image_uris.retrieve(
framework="sklearn",
region=region,
version="1.2-1",
instance_type="ml.m5.large",
image_scope="inference",
)

artifact_uri, key = _upload_raw_artifact(s3_client, bucket, prefix)
s3_keys.append(key)
src_dir = _make_source_code_dir()

model_builder = ModelBuilder(
image_uri=image_uri,
source_code=SourceCode(source_dir=src_dir, entry_script="inference.py"),
s3_model_data_url=artifact_uri,
role_arn=role,
instance_type="ml.m5.large",
sagemaker_session=session,
)
model_builder.model_path = f"/tmp/sagemaker/model-builder/{unique_id}"
os.makedirs(model_builder.model_path, exist_ok=True)
model_builder.dependencies = []

core_model = model_builder.build(
model_name=f"{MODEL_NAME_PREFIX}-{unique_id}", mode=Mode.SAGEMAKER_ENDPOINT
)

# A repack must have produced a new artifact (not the raw one)
repacked = model_builder.repacked_model_data
assert repacked is not None
assert repacked != artifact_uri
s3_keys.append(repacked.partition(f"s3://{bucket}/")[2])

# The repacked artifact must contain the inference code under code/
members = _tar_members(s3_client, repacked)
assert any(m.endswith("code/inference.py") for m in members), members
assert any(m.endswith("model.json") for m in members), members

# Script-mode env vars must be wired up on the model container
env = core_model.primary_container.environment or {}
assert env.get("SAGEMAKER_PROGRAM") == "inference.py"
assert env.get("SAGEMAKER_SUBMIT_DIRECTORY") == "/opt/ml/model/code"

finally:
if core_model is not None:
try:
core_model.delete()
except Exception: # noqa
pass
for k in s3_keys:
try:
s3_client.delete_object(Bucket=bucket, Key=k)
except Exception: # noqa
pass
31 changes: 27 additions & 4 deletions sagemaker-serve/tests/unit/test_model_builder_build.py
Original file line number Diff line number Diff line change
Expand Up @@ -386,10 +386,33 @@ def test_upload_code_local_mode(self, mock_determine):

self.assertIsNone(builder.uploaded_code)

@unittest.skip("Complex mocking required for repack_model with file system operations")
def test_upload_code_with_repack(self):
"""Test _upload_code with repacking."""
pass
@patch("sagemaker.serve.model_builder.repack_model")
@patch("sagemaker.core.workflow.is_pipeline_variable", return_value=False)
@patch("sagemaker.core.s3.determine_bucket_and_prefix")
def test_upload_code_repack_uses_script_dependencies(
self, mock_determine, mock_is_pipeline, mock_repack
):
"""_upload_code(repack=True) must pass script_dependencies (a list) to
repack_model, not the deprecated self.dependencies auto-detect dict."""
mock_determine.return_value = ("test-bucket", "test-prefix")

builder = ModelBuilder(
model=Mock(),
role_arn="arn:aws:iam::123456789012:role/TestRole",
sagemaker_session=self.mock_session,
)
builder.bucket = None
builder.entry_point = "inference.py"
builder.source_dir = "/path/to/code"
builder.script_dependencies = ["/path/to/requirements.txt"]
builder.dependencies = {"auto": True} # deprecated dict, must NOT be used
builder.s3_model_data_url = "s3://bucket/model.tar.gz"
builder.model_kms_key = None

builder._upload_code("test-prefix", repack=True)

mock_repack.assert_called_once()
assert mock_repack.call_args.kwargs["dependencies"] == ["/path/to/requirements.txt"]

@unittest.skip("Complex file system mocking - os.stat requires real file paths")
@patch('sagemaker.core.s3.determine_bucket_and_prefix')
Expand Down
114 changes: 114 additions & 0 deletions sagemaker-serve/tests/unit/test_model_builder_methods.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
from sagemaker.core.deserializers import JSONDeserializer, TorchTensorDeserializer
from sagemaker.serve.constants import Framework
from sagemaker.serve.mode.function_pointers import Mode
from sagemaker.core.training.configs import SourceCode


class TestModelBuilderSimpleMethods:
Expand Down Expand Up @@ -455,6 +456,119 @@ def test_build_for_passthrough_sets_s3_upload_path_none_when_no_model_path(
assert builder.s3_upload_path is None


class TestBuildForPassthroughSourceCodeRepack:
"""Tests for _build_for_passthrough() repacking source_code into the model artifact.

When source_code is supplied together with a model artifact, build() must repack
the code into the artifact instead of dropping it.
"""

def _make_mock_session(self):
mock_session = Mock()
mock_session.boto_region_name = "us-west-2"
mock_session.config = {}
mock_session.boto_session = Mock()
mock_session.boto_session.region_name = "us-west-2"
return mock_session

def _make_builder(self, **kwargs):
return ModelBuilder(
image_uri="123456789.dkr.ecr.us-west-2.amazonaws.com/my-image:latest",
mode=Mode.SAGEMAKER_ENDPOINT,
role_arn="arn:aws:iam::123456789012:role/TestRole",
source_code=SourceCode(source_dir="./code", entry_script="inference.py"),
sagemaker_session=self._make_mock_session(),
**kwargs,
)

@patch.object(ModelBuilder, "_create_model")
def test_repacks_and_preserves_script_mode_vars_with_s3_model_data_url(self, mock_create):
"""source_code + s3_model_data_url artifact -> repack path, script-mode vars kept."""
mock_create.return_value = Mock()

builder = self._make_builder(s3_model_data_url="s3://bucket/model.tar.gz")

builder._build_for_passthrough()

# script-mode vars must NOT be cleared (that would skip repack)
assert builder.entry_point == "inference.py"
assert builder.source_dir == "./code"
assert builder.s3_upload_path is None
mock_create.assert_called_once()

@patch.object(ModelBuilder, "_create_model")
def test_bridges_model_path_to_s3_model_data_url(self, mock_create):
"""When only model_path (S3) is given, it is bridged to s3_model_data_url for repack."""
mock_create.return_value = Mock()

builder = self._make_builder()
builder.model_path = "s3://bucket/model.tar.gz"

builder._build_for_passthrough()

assert builder.s3_model_data_url == "s3://bucket/model.tar.gz"
assert builder.entry_point == "inference.py"
assert builder.source_dir == "./code"
mock_create.assert_called_once()

@patch.object(ModelBuilder, "_create_model")
def test_repack_branch_makes_is_repack_true(self, mock_create):
"""After the repack branch runs, is_repack() must return True so
_prepare_container_def_base triggers _upload_code(repack=True)."""
mock_create.return_value = Mock()

builder = self._make_builder(s3_model_data_url="s3://bucket/model.tar.gz")

builder._build_for_passthrough()

assert builder.is_repack() is True

@patch.object(ModelBuilder, "_create_model")
def test_pure_image_passthrough_without_source_code_still_clears_vars(self, mock_create):
"""No source_code -> pure image passthrough clears script-mode vars (unchanged)."""
mock_create.return_value = Mock()

builder = ModelBuilder(
image_uri="123456789.dkr.ecr.us-west-2.amazonaws.com/my-image:latest",
mode=Mode.SAGEMAKER_ENDPOINT,
role_arn="arn:aws:iam::123456789012:role/TestRole",
s3_model_data_url="s3://bucket/model.tar.gz",
sagemaker_session=self._make_mock_session(),
)

builder._build_for_passthrough()

assert builder.entry_point is None
assert builder.source_dir is None
mock_create.assert_called_once()

@patch.object(ModelBuilder, "_create_model")
def test_entry_point_without_source_dir_warns_and_does_not_repack(self, mock_create):
"""entry_script without source_dir cannot be repacked: warn, don't repack."""
mock_create.return_value = Mock()

builder = ModelBuilder(
image_uri="123456789.dkr.ecr.us-west-2.amazonaws.com/my-image:latest",
mode=Mode.SAGEMAKER_ENDPOINT,
role_arn="arn:aws:iam::123456789012:role/TestRole",
source_code=SourceCode(entry_script="inference.py"),
s3_model_data_url="s3://bucket/model.tar.gz",
sagemaker_session=self._make_mock_session(),
)
# sanity: no source_dir -> repack branch is not eligible
assert builder.source_dir is None
assert builder.entry_point == "inference.py"

with patch("sagemaker.serve.model_builder.logger") as mock_logger:
builder._build_for_passthrough()
mock_logger.warning.assert_called_once()

# fall-through nulls script-mode vars (no repack occurred)
assert builder.entry_point is None
assert builder.source_dir is None
mock_create.assert_called_once()


class TestGetDockerClient:
"""Tests for _get_docker_client() Studio Docker client initialization.

Expand Down
Loading