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
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@ dependencies = [
train = ["sagemaker-train"]
serve = ["sagemaker-serve"]
mlops = ["sagemaker-mlops"]
torch = ["sagemaker-serve[torch]"]
all = ["sagemaker-train", "sagemaker-serve", "sagemaker-mlops"]

[project.urls]
Expand Down
3 changes: 2 additions & 1 deletion sagemaker-core/src/sagemaker/core/deserializers/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -368,7 +368,8 @@ def __init__(self, accept="tensor/pt"):
except ImportError as e:
raise ImportError(
"Unable to import torch. Please install torch to use TorchTensorDeserializer: "
"pip install 'sagemaker-core[torch]'"
"pip install 'sagemaker-core[torch]' "
"(or 'sagemaker-serve[torch]' if you installed sagemaker-serve)"
) from e

def deserialize(self, stream, content_type="tensor/pt"):
Expand Down
25 changes: 15 additions & 10 deletions sagemaker-core/src/sagemaker/core/serializers/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -443,17 +443,22 @@ class TorchTensorSerializer(SimpleBaseSerializer):

def __init__(self, content_type="tensor/pt"):
super(TorchTensorSerializer, self).__init__(content_type=content_type)
try:
from torch import Tensor
self.numpy_serializer = NumpySerializer()

self.torch_tensor = Tensor
except ImportError as e:
raise ImportError(
"Unable to import torch. Please install torch to use TorchTensorSerializer: "
"pip install 'sagemaker-core[torch]'"
) from e
@staticmethod
def _is_torch_tensor(data):
"""Recognize a torch.Tensor without importing torch.

self.numpy_serializer = NumpySerializer()
Serialization only needs the tensor's own detach()/numpy() methods, so
the type is identified structurally. A caller holding a real tensor
already has torch installed; importing it here would add nothing but a
multi-hundred-megabyte dependency for everyone else.
"""
return (
type(data).__module__.split(".")[0] == "torch"
and callable(getattr(data, "detach", None))
and callable(getattr(data, "numpy", None))
)

def serialize(self, data):
"""Serialize torch.Tensor to a buffer.
Expand All @@ -464,7 +469,7 @@ def serialize(self, data):
Returns:
raw-bytes: The data serialized as raw-bytes from the input.
"""
if isinstance(data, self.torch_tensor):
if self._is_torch_tensor(data):
try:
return self.numpy_serializer.serialize(data.detach().numpy())
except Exception as e:
Expand Down
22 changes: 18 additions & 4 deletions sagemaker-core/tests/unit/test_optional_torch_dependency.py
Original file line number Diff line number Diff line change
Expand Up @@ -121,16 +121,30 @@ def test_deserializer_module_imports_without_torch():
)


def test_torch_tensor_serializer_raises_import_error_without_torch():
"""Verify TorchTensorSerializer raises ImportError when torch is not installed."""
def test_torch_tensor_serializer_works_without_torch():
"""Serializing a tensor needs only its own detach()/numpy(), not the torch import,
so TorchTensorSerializer must construct and serialize without torch installed."""
import numpy as np

import sagemaker.core.serializers.base as ser_module

saved = {}
try:
saved = _block_torch()

with pytest.raises(ImportError, match="Unable to import torch"):
ser_module.TorchTensorSerializer()
class FakeTensor:
__module__ = "torch"

def detach(self):
return self

def numpy(self):
return np.array([1, 2, 3])

serializer = ser_module.TorchTensorSerializer()
assert serializer.serialize(FakeTensor())
with pytest.raises(ValueError, match="not a torch.Tensor"):
serializer.serialize([1, 2, 3])
finally:
_restore_torch(saved)

Expand Down
6 changes: 4 additions & 2 deletions sagemaker-serve/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -34,11 +34,13 @@ dependencies = [
"psutil",
"tritonclient[http]",
"onnx",
"onnxruntime",
"torch>=2.0.0"
"onnxruntime"
]

[project.optional-dependencies]
torch = [
"torch>=2.0.0",
]
test = [
"pytest",
"pytest-cov",
Expand Down
38 changes: 20 additions & 18 deletions sagemaker-serve/src/sagemaker/serve/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,17 +20,18 @@

Example:
Using Framework enum::

from sagemaker.serve.constants import Framework, DEFAULT_SERIALIZERS_BY_FRAMEWORK

# Get serializers for PyTorch
serializer, deserializer = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.PYTORCH]
serializer_cls, deserializer_cls = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.PYTORCH]
serializer, deserializer = serializer_cls(), deserializer_cls()
"""
from __future__ import absolute_import, annotations

# Standard library imports
from enum import Enum
from typing import Dict, Set, Tuple
from typing import Callable, Dict, Set, Tuple

# SageMaker imports
from sagemaker.serve.mode.function_pointers import Mode
Expand Down Expand Up @@ -96,7 +97,8 @@ class Framework(Enum):
Using framework enum::

if detected_framework == Framework.PYTORCH:
serializer, deserializer = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.PYTORCH]
ser_cls, deser_cls = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.PYTORCH]
serializer, deserializer = ser_cls(), deser_cls()
"""
XGBOOST = "XGBoost"
LDA = "LDA"
Expand All @@ -116,18 +118,18 @@ class Framework(Enum):
# Framework Serialization Mapping
# ========================================

DEFAULT_SERIALIZERS_BY_FRAMEWORK: Dict[Framework, Tuple] = {
Framework.XGBOOST: (LibSVMSerializer(), CSVDeserializer()),
Framework.LDA: (RecordSerializer(), RecordDeserializer()),
Framework.PYTORCH: (TorchTensorSerializer(), JSONDeserializer()),
Framework.TENSORFLOW: (NumpySerializer(), JSONDeserializer()),
Framework.MXNET: (RecordSerializer(), JSONDeserializer()),
Framework.CHAINER: (NumpySerializer(), JSONDeserializer()),
Framework.SKLEARN: (NumpySerializer(), NumpyDeserializer()),
Framework.HUGGINGFACE: (JSONSerializer(), JSONDeserializer()),
Framework.DJL: (JSONSerializer(), JSONDeserializer()),
Framework.SPARKML: (NumpySerializer(), JSONDeserializer()),
Framework.NTM: (RecordSerializer(), JSONDeserializer()),
Framework.SMD: (JSONSerializer(), JSONDeserializer()),
DEFAULT_SERIALIZERS_BY_FRAMEWORK: Dict[Framework, Tuple[Callable, Callable]] = {
Framework.XGBOOST: (LibSVMSerializer, CSVDeserializer),
Framework.LDA: (RecordSerializer, RecordDeserializer),
Framework.PYTORCH: (TorchTensorSerializer, JSONDeserializer),
Framework.TENSORFLOW: (NumpySerializer, JSONDeserializer),
Framework.MXNET: (RecordSerializer, JSONDeserializer),
Framework.CHAINER: (NumpySerializer, JSONDeserializer),
Framework.SKLEARN: (NumpySerializer, NumpyDeserializer),
Framework.HUGGINGFACE: (JSONSerializer, JSONDeserializer),
Framework.DJL: (JSONSerializer, JSONDeserializer),
Framework.SPARKML: (NumpySerializer, JSONDeserializer),
Framework.NTM: (RecordSerializer, JSONDeserializer),
Framework.SMD: (JSONSerializer, JSONDeserializer),
}

Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,13 @@ class TorchTensorTranslator:
"""Translate torch.Tensor from and to numpy.ndarray"""

def __init__(self) -> None:
import torch
try:
import torch
except ImportError as e:
raise ImportError(
"Unable to import torch. Translating to torch.Tensor requires torch: "
"pip install 'sagemaker-serve[torch]'"
) from e

self.convert_from_numpy = torch.from_numpy # pylint: disable=E1101
self.CONTENT_TYPE = "tensor/pt"
Expand Down
10 changes: 6 additions & 4 deletions sagemaker-serve/src/sagemaker/serve/model_builder_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1288,7 +1288,8 @@ def _fetch_serializer_and_deserializer_for_framework(self, framework: str) -> Tu
"""
framework_enum = self._normalize_framework_to_enum(framework)
if framework_enum and framework_enum in DEFAULT_SERIALIZERS_BY_FRAMEWORK:
return DEFAULT_SERIALIZERS_BY_FRAMEWORK[framework_enum]
serializer_cls, deserializer_cls = DEFAULT_SERIALIZERS_BY_FRAMEWORK[framework_enum]
return serializer_cls(), deserializer_cls()
return NumpySerializer(), JSONDeserializer()

def _normalize_framework_to_enum(
Expand Down Expand Up @@ -3036,9 +3037,10 @@ def _export_pytorch_to_onnx(

except ModuleNotFoundError:
raise ImportError(
"Launching Triton with ModelBuilder for a PyTorch model requires onnx module "
"but it was not found in your environment. "
"Checkout the instructions on the installation page of its repo: "
"Launching Triton with ModelBuilder for a PyTorch model requires the torch and "
"onnx modules but one of them was not found in your environment. "
"Install torch with: pip install 'sagemaker-serve[torch]'. "
"For onnx, check the instructions on the installation page of its repo: "
"https://onnxruntime.ai/docs/install/ "
"And follow the ones that match your environment. "
"Please note that you may need to restart your runtime after installation."
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,6 @@
import io
import logging
import threading
import torch
from typing import Optional, Type

from sagemaker.serve.spec.inference_spec import InferenceSpec
Expand Down Expand Up @@ -62,7 +61,12 @@ def __init__(
"Unable to import transformers, check if transformers is installed."
)

device = 0 if torch.cuda.is_available() else -1
try:
import torch

device = 0 if torch.cuda.is_available() else -1
except ImportError:
device = -1

self._load_model = pipeline(task, model=self.model, device=device)
except Exception:
Expand Down
24 changes: 12 additions & 12 deletions sagemaker-serve/tests/unit/test_constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,24 +56,24 @@ class TestDefaultSerializersByFramework(unittest.TestCase):
def test_all_frameworks_have_serializers(self):
for framework in Framework:
self.assertIn(framework, DEFAULT_SERIALIZERS_BY_FRAMEWORK)
serializer, deserializer = DEFAULT_SERIALIZERS_BY_FRAMEWORK[framework]
self.assertIsNotNone(serializer)
self.assertIsNotNone(deserializer)
serializer_cls, deserializer_cls = DEFAULT_SERIALIZERS_BY_FRAMEWORK[framework]
self.assertIsNotNone(serializer_cls)
self.assertIsNotNone(deserializer_cls)

def test_pytorch_serializers(self):
serializer, deserializer = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.PYTORCH]
self.assertEqual(serializer.__class__.__name__, "TorchTensorSerializer")
self.assertEqual(deserializer.__class__.__name__, "JSONDeserializer")
serializer_cls, deserializer_cls = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.PYTORCH]
self.assertEqual(serializer_cls.__name__, "TorchTensorSerializer")
self.assertEqual(deserializer_cls.__name__, "JSONDeserializer")

def test_tensorflow_serializers(self):
serializer, deserializer = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.TENSORFLOW]
self.assertEqual(serializer.__class__.__name__, "NumpySerializer")
self.assertEqual(deserializer.__class__.__name__, "JSONDeserializer")
serializer_cls, deserializer_cls = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.TENSORFLOW]
self.assertEqual(serializer_cls.__name__, "NumpySerializer")
self.assertEqual(deserializer_cls.__name__, "JSONDeserializer")

def test_sklearn_serializers(self):
serializer, deserializer = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.SKLEARN]
self.assertEqual(serializer.__class__.__name__, "NumpySerializer")
self.assertEqual(deserializer.__class__.__name__, "NumpyDeserializer")
serializer_cls, deserializer_cls = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.SKLEARN]
self.assertEqual(serializer_cls.__name__, "NumpySerializer")
self.assertEqual(deserializer_cls.__name__, "NumpyDeserializer")


if __name__ == "__main__":
Expand Down
9 changes: 6 additions & 3 deletions sagemaker-serve/tests/unit/test_model_builder_utils_new.py
Original file line number Diff line number Diff line change
Expand Up @@ -470,11 +470,14 @@ def test_fetch_serializer_for_known_framework(self, mock_default_serializers):
"""Test fetching serializer for known framework."""
mock_serializer = Mock()
mock_deserializer = Mock()
mock_default_serializers.__getitem__.return_value = (mock_serializer, mock_deserializer)
mock_default_serializers.__getitem__.return_value = (
Mock(return_value=mock_serializer),
Mock(return_value=mock_deserializer),
)
mock_default_serializers.__contains__.return_value = True

serializer, deserializer = self.utils._fetch_serializer_and_deserializer_for_framework("pytorch")

self.assertEqual(serializer, mock_serializer)
self.assertEqual(deserializer, mock_deserializer)

Expand Down
Loading
Loading