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
4 changes: 3 additions & 1 deletion sagemaker-serve/src/sagemaker/serve/model_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -3075,7 +3075,9 @@ def _build_single_modelbuilder(
self.sagemaker_session = (
sagemaker_session or self.sagemaker_session or self._create_session_with_region()
)
self.sagemaker_session.settings._local_download_dir = self.model_path
if isinstance(self.model_path, str) and not self.model_path.startswith("s3://"):
os.makedirs(self.model_path, exist_ok=True)
self.sagemaker_session.settings._local_download_dir = self.model_path

client = self.sagemaker_session.sagemaker_client
client._user_agent_creator.to_string = self._user_agent_decorator(
Expand Down
54 changes: 54 additions & 0 deletions sagemaker-serve/tests/unit/test_model_builder_core.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
from sagemaker.serve.spec.inference_spec import InferenceSpec
from sagemaker.train.model_trainer import ModelTrainer
from sagemaker.core.resources import TrainingJob, Model
from sagemaker.core.session_settings import SessionSettings
from sagemaker.core.training.configs import Compute, Networking, SourceCode


Expand Down Expand Up @@ -530,5 +531,58 @@ def test_prepare_for_mode_unsupported_mode(self):
self.assertIn("Unsupported deployment mode", str(context.exception))


class TestModelBuilderLocalDownloadDir(unittest.TestCase):
"""Test that build wires model_path into settings.local_download_dir correctly."""

def setUp(self):
"""Set up test fixtures."""
self.mock_session = Mock()
self.mock_session.boto_region_name = "us-west-2"
self.mock_session.settings = SessionSettings()

def _make_builder(self, model_path):
builder = ModelBuilder(
model=Mock(),
image_uri="123.dkr.ecr.us-west-2.amazonaws.com/custom:latest",
role_arn="arn:aws:iam::123456789012:role/TestRole",
sagemaker_session=self.mock_session,
)
builder.model_path = model_path
builder._passthrough = True
return builder

def _run(self, builder):
with patch.object(builder, "_get_serve_setting"), patch.object(
builder, "_is_model_customization", return_value=False
), patch.object(
builder, "_get_client_translators", return_value=(Mock(), Mock())
), patch.object(
builder, "_handle_mlflow_input"
), patch.object(
builder, "_build_validations"
), patch.object(
builder, "_build_for_passthrough", return_value=Mock()
):
builder._build_single_modelbuilder()

def test_local_model_path_is_created_and_used(self):
"""A local model_path is created on disk and set as local_download_dir."""
model_path = os.path.join(tempfile.mkdtemp(), "model-builder", "abc123")
self.assertFalse(os.path.exists(model_path))

builder = self._make_builder(model_path)
self._run(builder)

self.assertTrue(os.path.isdir(model_path))
self.assertEqual(self.mock_session.settings.local_download_dir, model_path)

def test_s3_model_path_leaves_local_download_dir_unset(self):
"""An s3:// model_path is not treated as a local download dir."""
builder = self._make_builder("s3://my-bucket/my-model/")
self._run(builder)

self.assertIsNone(self.mock_session.settings.local_download_dir)


if __name__ == "__main__":
unittest.main()
Loading