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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Add copy buttons to all
 blocks\n(function() {\n function addCopyButtons() {\n document.querySelectorAll('pre code').forEach(function(codeBlock) {\n if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;\n codeBlock.parentElement.setAttribute('data-copy-added', 'true');\n \n var btn = document.createElement('button');\n btn.textContent = 'Copy';\n btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';\n btn.onmouseover = function() { this.style.opacity = '1'; };\n btn.onmouseout = function() { this.style.opacity = '0.7'; };\n btn.onclick = function() {\n navigator.clipboard.writeText(codeBlock.textContent).then(function() {\n btn.textContent = 'Copied!';\n setTimeout(function() { btn.textContent = 'Copy'; }, 1500);\n });\n };\n codeBlock.parentElement.style.position = 'relative';\n codeBlock.parentElement.appendChild(btn);\n });\n }\n \n addCopyButtons();\n \n // Re-run on dynamic content\n var observer = new MutationObserver(addCopyButtons);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Add Copy Buttons to Code Blocks");
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Force GitHub README to respect dark mode\n(function() {\n var style = document.createElement('style');\n style.textContent = '\n .markdown-body {\n color-scheme: dark light;\n }\n .markdown-body pre { background: #161b22 !important; }\n .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; }\n .markdown-body table th, .markdown-body table td { border-color: #30363d !important; }\n .markdown-body img { background: #0d1117; }\n .markdown-body blockquote { border-left-color: #8b949e; }\n .markdown-body hr { border-color: #30363d; }\n ';\n document.head.appendChild(style);\n})();", "GitHub Dark Mode README Fix"); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Highlight search terms from Google/DuckDuckGo/Bing referrer\n(function() {\n var ref = document.referrer;\n var terms = [];\n \n if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) {\n var url = new URL(ref);\n var q = url.searchParams.get('q') || url.searchParams.get('p');\n if (q) {\n terms = q.split(/\\s+/).filter(function(t) { return t.length > 2; });\n }\n }\n \n if (terms.length === 0) return;\n \n var style = document.createElement('style');\n style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }';\n document.head.appendChild(style);\n \n function highlight(node) {\n if (node.nodeType === 3) { // text node\n var text = node.textContent;\n var found = false;\n terms.forEach(function(term) {\n var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\\]\\\\]/g, '\\\\') + ')', 'gi');\n if (regex.test(text)) {\n found = true;\n var frag = document.createDocumentFragment();\n var parts = text.split(regex);\n parts.forEach(function(part, i) {\n if (i % 2 === 0) {\n frag.appendChild(document.createTextNode(part));\n } else {\n var span = document.createElement('span');\n span.className = 'userscript-highlight';\n span.textContent = part;\n frag.appendChild(span);\n }\n });\n node.parentNode.replaceChild(frag, node);\n }\n });\n } else if (node.nodeType === 1 && node.childNodes) { // element\n var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT'];\n if (!skipTags.includes(node.tagName)) {\n Array.from(node.childNodes).forEach(highlight);\n }\n }\n }\n \n highlight(document.body);\n \n // Re-highlight on dynamic content\n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1 || node.nodeType === 3) highlight(node);\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Highlight Search Terms"); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Strip utm_, fbclid, gclid, etc. from all links on page\n(function() {\n var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content',\n 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid',\n 'ref', 'ref_src', 'source', 'medium', 'campaign'];\n \n function cleanUrl(url) {\n try {\n var u = new URL(url, window.location.origin);\n var changed = false;\n trackingParams.forEach(function(p) {\n if (u.searchParams.has(p)) {\n u.searchParams.delete(p);\n changed = true;\n }\n });\n return changed ? u.toString() : url;\n } catch (e) {\n return url;\n }\n }\n \n function cleanLinks() {\n document.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n \n cleanLinks();\n \n var observer = new MutationObserver(function(mutations) {\n mutations.forEach(function(m) {\n m.addedNodes.forEach(function(node) {\n if (node.nodeType === 1) {\n if (node.tagName === 'A') cleanLinks();\n node.querySelectorAll('a[href]').forEach(function(a) {\n var clean = cleanUrl(a.href);\n if (clean !== a.href) a.href = clean;\n });\n }\n });\n });\n });\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "Remove Tracking Parameters from Links"); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + '
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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Auto-enable theater mode on YouTube\n(function() {\n function tryTheater() {\n var btn = document.querySelector('button[aria-label=\"Theater mode\"], ytd-player #player button[title=\"Theater mode\"]');\n if (btn && !btn.classList.contains('activated')) {\n btn.click();\n }\n }\n \n // Try immediately\n tryTheater();\n \n // Try after navigation (SPA)\n var lastUrl = location.href;\n setInterval(function() {\n if (location.href !== lastUrl) {\n lastUrl = location.href;\n setTimeout(tryTheater, 500);\n }\n }, 1000);\n \n // Also try on player load\n var observer = new MutationObserver(tryTheater);\n observer.observe(document.body, { childList: true, subtree: true });\n})();", "YouTube Theater Mode Default"); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Remove or un-stick sticky/fixed headers that block content\n(function() {\n function unstick() {\n document.querySelectorAll('header, nav, [role=\"banner\"], .header, .navbar, .sticky, .fixed-top, [style*=\"position: fixed\"], [style*=\"position:sticky\"]').forEach(function(el) {\n if (el.style.position === 'fixed' || el.style.position === 'sticky' || \n getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') {\n el.style.position = 'static';\n el.style.top = 'auto';\n el.style.zIndex = 'auto';\n }\n });\n }\n \n unstick();\n \n var observer = new MutationObserver(unstick);\n observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] });\n})();", "Kill Sticky Headers"); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + '
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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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
, 'i'); if (__m === '*' || __re.test(location.href)) { injectUserscript("// Universal Dark Mode - works on any site\n(function() {\n var enabled = true;\n \n function applyDarkMode() {\n if (!enabled) return;\n \n // Create style element if it doesn't exist\n var style = document.getElementById('universal-dark-mode-style');\n if (!style) {\n style = document.createElement('style');\n style.id = 'universal-dark-mode-style';\n document.head.appendChild(style);\n }\n \n // Dark mode CSS - inverts colors but preserves images/video\n style.textContent = '\n /* Invert everything except media */\n html {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #1a1a2e !important;\n }\n \n /* Restore images, videos, iframes, canvas */\n img, video, iframe, canvas, svg, picture, [style*=\"background-image\"] {\n filter: invert(1) hue-rotate(180deg) !important;\n }\n \n /* Preserve specific elements that should not be inverted */\n .no-dark-mode, .no-dark-mode *,\n [data-theme=\"light\"], [data-theme=\"light\"],\n .ace_editor, .ace_editor *,\n .CodeMirror, .CodeMirror *,\n .monaco-editor, .monaco-editor *,\n .markdown-body pre, .markdown-body pre *,\n .highlight, .highlight *,\n pre code, pre code * {\n filter: none !important;\n }\n \n /* Fix common UI elements */\n .modal, .popup, .dropdown-menu, .tooltip, .popover {\n filter: invert(1) hue-rotate(180deg) !important;\n background: #2d2d44 !important;\n border-color: #444 !important;\n }\n \n /* Scrollbars */\n ::-webkit-scrollbar { background: #1a1a2e !important; }\n ::-webkit-scrollbar-thumb { background: #444 !important; }\n ::-webkit-scrollbar-thumb:hover { background: #555 !important; }\n \n /* Selection */\n ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; }\n ';\n }\n \n function removeDarkMode() {\n var style = document.getElementById('universal-dark-mode-style');\n if (style) style.remove();\n }\n \n // Toggle with Alt+Shift+D\n document.addEventListener('keydown', function(e) {\n if (e.altKey && e.shiftKey && e.key === 'D') {\n e.preventDefault();\n enabled = !enabled;\n if (enabled) {\n applyDarkMode();\n console.log('[Universal Dark Mode] Enabled');\n } else {\n removeDarkMode();\n console.log('[Universal Dark Mode] Disabled');\n }\n }\n });\n \n // Apply on load\n applyDarkMode();\n \n // Re-apply on dynamic content\n var observer = new MutationObserver(function(mutations) {\n if (enabled && !document.getElementById('universal-dark-mode-style')) {\n applyDarkMode();\n }\n });\n observer.observe(document.head, { childList: true });\n \n console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle');\n})();", "Universal Dark Mode"); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })();
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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 DownExpand 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 numberDiff line numberDiff 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 numberDiff line numberDiff 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