Skip to content
Open
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
8 changes: 7 additions & 1 deletion src/sagemaker/predictor_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,9 @@ def __init__(
predictor (sagemaker.predictor.Predictor): General ``Predictor``
object has useful methods and variables. ``AsyncPredictor``
stands on top of it with capability for async inference.
name (str): Optional. Name used as the prefix of the Amazon S3 key when
input data is uploaded for async inference. If not specified, the
endpoint name is used. (Default: None)
"""
warn_v2_deprecation(
feature="AsyncPredictor",
Expand Down Expand Up @@ -174,10 +177,13 @@ def _upload_data_to_s3(
my_uuid = str(uuid.uuid4())
timestamp = sagemaker_timestamp()
bucket = self.sagemaker_session.default_bucket()
# ``name`` is optional; fall back to the endpoint name so the default
# ``AsyncPredictor(predictor)`` construction can upload input data.
base_name = self.name or self.endpoint_name
key = s3.s3_path_join(
self.sagemaker_session.default_bucket_prefix,
"async-endpoint-inputs",
name_from_base(self.name, short=True),
name_from_base(base_name, short=True),
"{}-{}".format(timestamp, my_uuid),
)

Expand Down
25 changes: 25 additions & 0 deletions tests/unit/test_predictor_async.py
Original file line number Diff line number Diff line change
Expand Up @@ -170,6 +170,31 @@ def test_async_predict_call_with_data():
assert result.output_path == ASYNC_OUTPUT_LOCATION


def test_async_predict_call_with_data_and_no_name():
# Regression test for https://github.com/aws/sagemaker-python-sdk/issues/3210
sagemaker_session = empty_sagemaker_session()
predictor_async = AsyncPredictor(Predictor(ENDPOINT, sagemaker_session))
assert predictor_async.name is None

result = predictor_async.predict_async(data=DUMMY_DATA)

_, put_kwargs = sagemaker_session.s3_client.put_object.call_args
assert put_kwargs["Bucket"] == BUCKET_NAME
assert put_kwargs["Key"].startswith("async-endpoint-inputs/{}-".format(ENDPOINT))
assert predictor_async._input_path == "s3://{}/{}".format(BUCKET_NAME, put_kwargs["Key"])
assert result.output_path == ASYNC_OUTPUT_LOCATION


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

predictor_async.predict_async(data=DUMMY_DATA)

_, put_kwargs = sagemaker_session.s3_client.put_object.call_args
assert put_kwargs["Key"].startswith("async-endpoint-inputs/{}-".format(ASYNC_PREDICTOR))


def test_async_predict_call_with_data_and_input_path():
sagemaker_session = empty_sagemaker_session()
predictor_async = AsyncPredictor(Predictor(ENDPOINT, sagemaker_session))
Expand Down
Loading