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 src/sagemaker/predictor_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,14 +57,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
63 changes: 63 additions & 0 deletions tests/unit/test_predictor_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,11 +12,15 @@
# language governing permissions and limitations under the License.
from __future__ import absolute_import

import io

import pytest
from mock import Mock
from sagemaker.async_inference.waiter_config import WaiterConfig
from sagemaker.deserializers import JSONDeserializer
from sagemaker.predictor import Predictor
from sagemaker.predictor_async import AsyncPredictor
from sagemaker.serializers import JSONSerializer
from sagemaker.exceptions import AsyncInferenceModelError, PollingTimeoutError

ENDPOINT = "mxnet_endpoint"
Expand Down Expand Up @@ -494,3 +498,62 @@ def test_list_monitors():

predictor_async.list_monitors()
predictor.list_monitors.assert_called_with()


def _sagemaker_session_with_json_output():
sagemaker_session = empty_sagemaker_session()
sagemaker_session.s3_client.get_waiter = Mock(name="get_waiter")
sagemaker_session.s3_client.get_object = Mock(
name="get_object",
return_value={
"Body": io.BytesIO(b'{"result": [1, 2, 3]}'),
"ContentType": "application/json",
},
)
return sagemaker_session


def test_async_predictor_deserializer_override_is_used_for_accept_and_result():
sagemaker_session = _sagemaker_session_with_json_output()
predictor_async = AsyncPredictor(Predictor(ENDPOINT, sagemaker_session))

predictor_async.deserializer = JSONDeserializer()

assert isinstance(predictor_async.predictor.deserializer, JSONDeserializer)

result = predictor_async.predict(
input_path=ASYNC_INPUT_LOCATION, waiter_config=DEFAULT_WAITER_CONFIG
)

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


def test_async_predictor_serializer_override_is_used_for_upload():
sagemaker_session = empty_sagemaker_session()
predictor_async = AsyncPredictor(Predictor(ENDPOINT, sagemaker_session))
predictor_async.name = ASYNC_PREDICTOR

predictor_async.serializer = JSONSerializer()

assert isinstance(predictor_async.predictor.serializer, JSONSerializer)

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

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


def test_async_predictor_reflects_serializers_set_on_wrapped_predictor():
predictor = Predictor(ENDPOINT, empty_sagemaker_session())
predictor_async = AsyncPredictor(predictor)

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

assert predictor_async.serializer is serializer
assert predictor_async.deserializer is deserializer
Loading