Skip to content

Commit bf0f6de

Browse files
authored
Merge branch 'aws:master' into list-hyperparameters
2 parents 1b4e4a8 + 748910e commit bf0f6de

2 files changed

Lines changed: 57 additions & 1 deletion

File tree

sagemaker-serve/src/sagemaker/serve/model_builder.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -3075,7 +3075,9 @@ def _build_single_modelbuilder(
30753075
self.sagemaker_session = (
30763076
sagemaker_session or self.sagemaker_session or self._create_session_with_region()
30773077
)
3078-
self.sagemaker_session.settings._local_download_dir = self.model_path
3078+
if isinstance(self.model_path, str) and not self.model_path.startswith("s3://"):
3079+
os.makedirs(self.model_path, exist_ok=True)
3080+
self.sagemaker_session.settings._local_download_dir = self.model_path
30793081

30803082
client = self.sagemaker_session.sagemaker_client
30813083
client._user_agent_creator.to_string = self._user_agent_decorator(

sagemaker-serve/tests/unit/test_model_builder_core.py

Lines changed: 54 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@
1414
from sagemaker.serve.spec.inference_spec import InferenceSpec
1515
from sagemaker.train.model_trainer import ModelTrainer
1616
from sagemaker.core.resources import TrainingJob, Model
17+
from sagemaker.core.session_settings import SessionSettings
1718
from sagemaker.core.training.configs import Compute, Networking, SourceCode
1819

1920

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

532533

534+
class TestModelBuilderLocalDownloadDir(unittest.TestCase):
535+
"""Test that build wires model_path into settings.local_download_dir correctly."""
536+
537+
def setUp(self):
538+
"""Set up test fixtures."""
539+
self.mock_session = Mock()
540+
self.mock_session.boto_region_name = "us-west-2"
541+
self.mock_session.settings = SessionSettings()
542+
543+
def _make_builder(self, model_path):
544+
builder = ModelBuilder(
545+
model=Mock(),
546+
image_uri="123.dkr.ecr.us-west-2.amazonaws.com/custom:latest",
547+
role_arn="arn:aws:iam::123456789012:role/TestRole",
548+
sagemaker_session=self.mock_session,
549+
)
550+
builder.model_path = model_path
551+
builder._passthrough = True
552+
return builder
553+
554+
def _run(self, builder):
555+
with patch.object(builder, "_get_serve_setting"), patch.object(
556+
builder, "_is_model_customization", return_value=False
557+
), patch.object(
558+
builder, "_get_client_translators", return_value=(Mock(), Mock())
559+
), patch.object(
560+
builder, "_handle_mlflow_input"
561+
), patch.object(
562+
builder, "_build_validations"
563+
), patch.object(
564+
builder, "_build_for_passthrough", return_value=Mock()
565+
):
566+
builder._build_single_modelbuilder()
567+
568+
def test_local_model_path_is_created_and_used(self):
569+
"""A local model_path is created on disk and set as local_download_dir."""
570+
model_path = os.path.join(tempfile.mkdtemp(), "model-builder", "abc123")
571+
self.assertFalse(os.path.exists(model_path))
572+
573+
builder = self._make_builder(model_path)
574+
self._run(builder)
575+
576+
self.assertTrue(os.path.isdir(model_path))
577+
self.assertEqual(self.mock_session.settings.local_download_dir, model_path)
578+
579+
def test_s3_model_path_leaves_local_download_dir_unset(self):
580+
"""An s3:// model_path is not treated as a local download dir."""
581+
builder = self._make_builder("s3://my-bucket/my-model/")
582+
self._run(builder)
583+
584+
self.assertIsNone(self.mock_session.settings.local_download_dir)
585+
586+
533587
if __name__ == "__main__":
534588
unittest.main()

0 commit comments

Comments
 (0)