diff --git a/sagemaker-serve/src/sagemaker/serve/model_builder.py b/sagemaker-serve/src/sagemaker/serve/model_builder.py index 0d1e6e4746..c3d96c6399 100644 --- a/sagemaker-serve/src/sagemaker/serve/model_builder.py +++ b/sagemaker-serve/src/sagemaker/serve/model_builder.py @@ -4230,8 +4230,8 @@ def _deploy_for_ic(self, ic_data: Dict[str, Any], endpoint_name: str, **kwargs) endpoint_type=EndpointType.INFERENCE_COMPONENT_BASED, resources=resource_requirements, inference_component_name=ic_name, - instance_type=kwargs.get("instance_type", self.instance_type), - initial_instance_count=kwargs.get("initial_instance_count", 1), + instance_type=kwargs.pop("instance_type", self.instance_type), + initial_instance_count=kwargs.pop("initial_instance_count", 1), **kwargs, ) diff --git a/sagemaker-serve/tests/unit/test_model_builder_deploy.py b/sagemaker-serve/tests/unit/test_model_builder_deploy.py index 3a0fca3d8e..1a6d0817ba 100644 --- a/sagemaker-serve/tests/unit/test_model_builder_deploy.py +++ b/sagemaker-serve/tests/unit/test_model_builder_deploy.py @@ -521,6 +521,74 @@ def test_does_ic_exist_false(self): self.assertFalse(result) + def test_deploy_for_ic_creates_new_ic_with_explicit_instance_kwargs(self): + """Test _deploy_for_ic forwards instance_type and initial_instance_count to _deploy once.""" + builder = ModelBuilder( + model=Mock(), + role_arn="arn:aws:iam::123456789012:role/TestRole", + sagemaker_session=self.mock_session, + ) + builder._does_ic_exist = Mock(return_value=False) + builder._deploy = Mock(return_value="endpoint") + built_model = Mock() + resource_requirements = ResourceRequirements(requests={"memory": 1024, "copies": 1}) + ic_data = { + "Name": "test-ic", + "ResourceRequirements": resource_requirements, + "Model": built_model, + } + + result = builder._deploy_for_ic( + ic_data=ic_data, + endpoint_name="test-endpoint", + container_timeout_in_seconds=600, + instance_type="ml.g5.xlarge", + initial_instance_count=2, + ) + + self.assertEqual(result, "endpoint") + builder._deploy.assert_called_once_with( + built_model=built_model, + endpoint_name="test-endpoint", + endpoint_type=EndpointType.INFERENCE_COMPONENT_BASED, + resources=resource_requirements, + inference_component_name="test-ic", + instance_type="ml.g5.xlarge", + initial_instance_count=2, + container_timeout_in_seconds=600, + ) + + def test_deploy_for_ic_creates_new_ic_with_default_instance_kwargs(self): + """Test _deploy_for_ic falls back to builder instance_type and one instance.""" + builder = ModelBuilder( + model=Mock(), + role_arn="arn:aws:iam::123456789012:role/TestRole", + sagemaker_session=self.mock_session, + ) + builder.instance_type = "ml.c5.xlarge" + builder._does_ic_exist = Mock(return_value=False) + builder._deploy = Mock(return_value="endpoint") + built_model = Mock() + resource_requirements = ResourceRequirements(requests={"memory": 1024, "copies": 1}) + ic_data = { + "Name": "test-ic", + "ResourceRequirements": resource_requirements, + "Model": built_model, + } + + result = builder._deploy_for_ic(ic_data=ic_data, endpoint_name="test-endpoint") + + self.assertEqual(result, "endpoint") + builder._deploy.assert_called_once_with( + built_model=built_model, + endpoint_name="test-endpoint", + endpoint_type=EndpointType.INFERENCE_COMPONENT_BASED, + resources=resource_requirements, + inference_component_name="test-ic", + instance_type="ml.c5.xlarge", + initial_instance_count=1, + ) + class TestModelBuilderResetState(unittest.TestCase): """Test ModelBuilder _reset_build_state method."""