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
30 changes: 28 additions & 2 deletions sagemaker-serve/src/sagemaker/serve/predictor_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,14 +52,40 @@ def __init__(
else:
self.s3_client = self.sagemaker_session.s3_client

self.serializer = predictor.serializer
self.deserializer = predictor.deserializer
self.name = name
self._endpoint_config_name = None
self._model_names = None
self._context = None
self._input_path = None

@property
def serializer(self):
"""The serializer used to encode request data uploaded to Amazon S3.

Reads from and writes to the wrapped predictor, so it always stays in
sync with the serializer the underlying predictor uses.
"""
return self.predictor.serializer

@serializer.setter
def serializer(self, serializer):
"""Set the serializer on the wrapped predictor."""
self.predictor.serializer = serializer

@property
def deserializer(self):
"""The deserializer used to decode the async inference result.

Reads from and writes to the wrapped predictor, so the ``Accept`` header
and the decoding of the Amazon S3 output both honor the configured value.
"""
return self.predictor.deserializer

@deserializer.setter
def deserializer(self, deserializer):
"""Set the deserializer on the wrapped predictor."""
self.predictor.deserializer = deserializer

def predict(
self,
data=None,
Expand Down
77 changes: 77 additions & 0 deletions sagemaker-serve/tests/unit/test_predictor_async.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,9 @@
import io
import unittest
from unittest.mock import Mock, patch

from sagemaker.core.deserializers import JSONDeserializer
from sagemaker.core.serializers import JSONSerializer
from sagemaker.serve.predictor_async import AsyncPredictor


Expand Down Expand Up @@ -107,5 +111,78 @@ def test_disable_data_capture(self):
self.mock_predictor.disable_data_capture.assert_called_once()


class _StubPredictor:
"""Minimal predictor exposing the attributes ``AsyncPredictor`` relies on."""

def __init__(self, sagemaker_session, serializer, deserializer):
self.endpoint_name = "test-endpoint"
self.sagemaker_session = sagemaker_session
self.serializer = serializer
self.deserializer = deserializer

@property
def accept(self):
return self.deserializer.ACCEPT

def _handle_response(self, response):
return self.deserializer.deserialize(response["Body"], response["ContentType"])


class TestAsyncPredictorSerializerOverrides(unittest.TestCase):
"""Overrides set on AsyncPredictor must reach the request and the result (issue #3100)."""

def setUp(self):
self.sagemaker_session = Mock()
self.sagemaker_session.default_bucket.return_value = "bucket"
self.sagemaker_session.default_bucket_prefix = None
self.sagemaker_session.sagemaker_runtime_client.invoke_endpoint_async.return_value = {
"OutputLocation": "s3://bucket/output",
}
self.sagemaker_session.s3_client.get_object.return_value = {
"Body": io.BytesIO(b'{"result": [1, 2, 3]}'),
"ContentType": "application/json",
}
default_deserializer = Mock()
default_deserializer.ACCEPT = ("*/*",)
self.predictor = _StubPredictor(self.sagemaker_session, Mock(), default_deserializer)

def test_deserializer_override_is_used_for_accept_and_result(self):
async_predictor = AsyncPredictor(self.predictor)

async_predictor.deserializer = JSONDeserializer()

self.assertIsInstance(self.predictor.deserializer, JSONDeserializer)

result = async_predictor.predict(input_path="s3://bucket/input")

_, kwargs = self.sagemaker_session.sagemaker_runtime_client.invoke_endpoint_async.call_args
self.assertEqual(kwargs["Accept"], "application/json")
self.assertEqual(result, {"result": [1, 2, 3]})

def test_serializer_override_is_used_for_upload(self):
async_predictor = AsyncPredictor(self.predictor, name="test")

async_predictor.serializer = JSONSerializer()

self.assertIsInstance(self.predictor.serializer, JSONSerializer)

async_predictor.predict_async(data={"hi": "there"})

_, kwargs = self.sagemaker_session.s3_client.put_object.call_args
self.assertEqual(kwargs["Body"], '{"hi": "there"}')
self.assertEqual(kwargs["ContentType"], "application/json")

def test_reflects_serializers_set_on_wrapped_predictor(self):
async_predictor = AsyncPredictor(self.predictor)

serializer = JSONSerializer()
deserializer = JSONDeserializer()
self.predictor.serializer = serializer
self.predictor.deserializer = deserializer

self.assertIs(async_predictor.serializer, serializer)
self.assertIs(async_predictor.deserializer, deserializer)


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