Skip to content

Commit 1725779

Browse files
committed
fix(serve): make torch an optional dependency
DEFAULT_SERIALIZERS_BY_FRAMEWORK instantiated TorchTensorSerializer at module scope, so importing sagemaker.serve required torch. Store the classes instead and instantiate on lookup, and move torch to an extra. Fixes #5531
1 parent e698fc6 commit 1725779

6 files changed

Lines changed: 142 additions & 36 deletions

File tree

sagemaker-serve/pyproject.toml

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -34,11 +34,13 @@ dependencies = [
3434
"psutil",
3535
"tritonclient[http]",
3636
"onnx",
37-
"onnxruntime",
38-
"torch>=2.0.0"
37+
"onnxruntime"
3938
]
4039

4140
[project.optional-dependencies]
41+
torch = [
42+
"torch>=2.0.0",
43+
]
4244
test = [
4345
"pytest",
4446
"pytest-cov",

sagemaker-serve/src/sagemaker/serve/constants.py

Lines changed: 22 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -20,17 +20,18 @@
2020
2121
Example:
2222
Using Framework enum::
23-
23+
2424
from sagemaker.serve.constants import Framework, DEFAULT_SERIALIZERS_BY_FRAMEWORK
25-
25+
2626
# Get serializers for PyTorch
27-
serializer, deserializer = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.PYTORCH]
27+
serializer_cls, deserializer_cls = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.PYTORCH]
28+
serializer, deserializer = serializer_cls(), deserializer_cls()
2829
"""
2930
from __future__ import absolute_import, annotations
3031

3132
# Standard library imports
3233
from enum import Enum
33-
from typing import Dict, Set, Tuple
34+
from typing import Callable, Dict, Set, Tuple
3435

3536
# SageMaker imports
3637
from sagemaker.serve.mode.function_pointers import Mode
@@ -96,7 +97,8 @@ class Framework(Enum):
9697
Using framework enum::
9798
9899
if detected_framework == Framework.PYTORCH:
99-
serializer, deserializer = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.PYTORCH]
100+
ser_cls, deser_cls = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.PYTORCH]
101+
serializer, deserializer = ser_cls(), deser_cls()
100102
"""
101103
XGBOOST = "XGBoost"
102104
LDA = "LDA"
@@ -116,18 +118,20 @@ class Framework(Enum):
116118
# Framework Serialization Mapping
117119
# ========================================
118120

119-
DEFAULT_SERIALIZERS_BY_FRAMEWORK: Dict[Framework, Tuple] = {
120-
Framework.XGBOOST: (LibSVMSerializer(), CSVDeserializer()),
121-
Framework.LDA: (RecordSerializer(), RecordDeserializer()),
122-
Framework.PYTORCH: (TorchTensorSerializer(), JSONDeserializer()),
123-
Framework.TENSORFLOW: (NumpySerializer(), JSONDeserializer()),
124-
Framework.MXNET: (RecordSerializer(), JSONDeserializer()),
125-
Framework.CHAINER: (NumpySerializer(), JSONDeserializer()),
126-
Framework.SKLEARN: (NumpySerializer(), NumpyDeserializer()),
127-
Framework.HUGGINGFACE: (JSONSerializer(), JSONDeserializer()),
128-
Framework.DJL: (JSONSerializer(), JSONDeserializer()),
129-
Framework.SPARKML: (NumpySerializer(), JSONDeserializer()),
130-
Framework.NTM: (RecordSerializer(), JSONDeserializer()),
131-
Framework.SMD: (JSONSerializer(), JSONDeserializer()),
121+
# Values are classes, not instances: instantiating TorchTensorSerializer at
122+
# import time forces `import torch` on every consumer of this module.
123+
DEFAULT_SERIALIZERS_BY_FRAMEWORK: Dict[Framework, Tuple[Callable, Callable]] = {
124+
Framework.XGBOOST: (LibSVMSerializer, CSVDeserializer),
125+
Framework.LDA: (RecordSerializer, RecordDeserializer),
126+
Framework.PYTORCH: (TorchTensorSerializer, JSONDeserializer),
127+
Framework.TENSORFLOW: (NumpySerializer, JSONDeserializer),
128+
Framework.MXNET: (RecordSerializer, JSONDeserializer),
129+
Framework.CHAINER: (NumpySerializer, JSONDeserializer),
130+
Framework.SKLEARN: (NumpySerializer, NumpyDeserializer),
131+
Framework.HUGGINGFACE: (JSONSerializer, JSONDeserializer),
132+
Framework.DJL: (JSONSerializer, JSONDeserializer),
133+
Framework.SPARKML: (NumpySerializer, JSONDeserializer),
134+
Framework.NTM: (RecordSerializer, JSONDeserializer),
135+
Framework.SMD: (JSONSerializer, JSONDeserializer),
132136
}
133137

sagemaker-serve/src/sagemaker/serve/model_builder_utils.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1288,7 +1288,8 @@ def _fetch_serializer_and_deserializer_for_framework(self, framework: str) -> Tu
12881288
"""
12891289
framework_enum = self._normalize_framework_to_enum(framework)
12901290
if framework_enum and framework_enum in DEFAULT_SERIALIZERS_BY_FRAMEWORK:
1291-
return DEFAULT_SERIALIZERS_BY_FRAMEWORK[framework_enum]
1291+
serializer_cls, deserializer_cls = DEFAULT_SERIALIZERS_BY_FRAMEWORK[framework_enum]
1292+
return serializer_cls(), deserializer_cls()
12921293
return NumpySerializer(), JSONDeserializer()
12931294

12941295
def _normalize_framework_to_enum(

sagemaker-serve/tests/unit/test_constants.py

Lines changed: 12 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -56,24 +56,24 @@ class TestDefaultSerializersByFramework(unittest.TestCase):
5656
def test_all_frameworks_have_serializers(self):
5757
for framework in Framework:
5858
self.assertIn(framework, DEFAULT_SERIALIZERS_BY_FRAMEWORK)
59-
serializer, deserializer = DEFAULT_SERIALIZERS_BY_FRAMEWORK[framework]
60-
self.assertIsNotNone(serializer)
61-
self.assertIsNotNone(deserializer)
59+
serializer_cls, deserializer_cls = DEFAULT_SERIALIZERS_BY_FRAMEWORK[framework]
60+
self.assertIsNotNone(serializer_cls)
61+
self.assertIsNotNone(deserializer_cls)
6262

6363
def test_pytorch_serializers(self):
64-
serializer, deserializer = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.PYTORCH]
65-
self.assertEqual(serializer.__class__.__name__, "TorchTensorSerializer")
66-
self.assertEqual(deserializer.__class__.__name__, "JSONDeserializer")
64+
serializer_cls, deserializer_cls = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.PYTORCH]
65+
self.assertEqual(serializer_cls.__name__, "TorchTensorSerializer")
66+
self.assertEqual(deserializer_cls.__name__, "JSONDeserializer")
6767

6868
def test_tensorflow_serializers(self):
69-
serializer, deserializer = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.TENSORFLOW]
70-
self.assertEqual(serializer.__class__.__name__, "NumpySerializer")
71-
self.assertEqual(deserializer.__class__.__name__, "JSONDeserializer")
69+
serializer_cls, deserializer_cls = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.TENSORFLOW]
70+
self.assertEqual(serializer_cls.__name__, "NumpySerializer")
71+
self.assertEqual(deserializer_cls.__name__, "JSONDeserializer")
7272

7373
def test_sklearn_serializers(self):
74-
serializer, deserializer = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.SKLEARN]
75-
self.assertEqual(serializer.__class__.__name__, "NumpySerializer")
76-
self.assertEqual(deserializer.__class__.__name__, "NumpyDeserializer")
74+
serializer_cls, deserializer_cls = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.SKLEARN]
75+
self.assertEqual(serializer_cls.__name__, "NumpySerializer")
76+
self.assertEqual(deserializer_cls.__name__, "NumpyDeserializer")
7777

7878

7979
if __name__ == "__main__":

sagemaker-serve/tests/unit/test_model_builder_utils_new.py

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -470,11 +470,14 @@ def test_fetch_serializer_for_known_framework(self, mock_default_serializers):
470470
"""Test fetching serializer for known framework."""
471471
mock_serializer = Mock()
472472
mock_deserializer = Mock()
473-
mock_default_serializers.__getitem__.return_value = (mock_serializer, mock_deserializer)
473+
mock_default_serializers.__getitem__.return_value = (
474+
Mock(return_value=mock_serializer),
475+
Mock(return_value=mock_deserializer),
476+
)
474477
mock_default_serializers.__contains__.return_value = True
475-
478+
476479
serializer, deserializer = self.utils._fetch_serializer_and_deserializer_for_framework("pytorch")
477-
480+
478481
self.assertEqual(serializer, mock_serializer)
479482
self.assertEqual(deserializer, mock_deserializer)
480483

Lines changed: 96 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,96 @@
1+
# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved.
2+
#
3+
# Licensed under the Apache License, Version 2.0 (the "License"). You
4+
# may not use this file except in compliance with the License. A copy of
5+
# the License is located at
6+
#
7+
# http://aws.amazon.com/apache2.0/
8+
#
9+
# or in the "license" file accompanying this file. This file is
10+
# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF
11+
# ANY KIND, either express or implied. See the License for the specific
12+
# language governing permissions and limitations under the License.
13+
"""Tests to verify torch dependency is optional in sagemaker-serve.
14+
15+
Runs in a subprocess so that blocking torch cannot affect other tests in
16+
the session.
17+
"""
18+
from __future__ import absolute_import
19+
20+
import subprocess
21+
import sys
22+
import textwrap
23+
24+
25+
def _run_without_torch(body):
26+
"""Execute body in a subprocess where importing torch raises ImportError."""
27+
script = textwrap.dedent(
28+
"""
29+
import sys
30+
31+
class _TorchBlocker:
32+
\"\"\"Meta path finder that makes `import torch` fail.\"\"\"
33+
34+
def find_spec(self, name, path=None, target=None):
35+
if name == "torch" or name.startswith("torch."):
36+
raise ImportError("torch is blocked for this test")
37+
return None
38+
39+
sys.meta_path.insert(0, _TorchBlocker())
40+
for _name in [n for n in sys.modules if n == "torch" or n.startswith("torch.")]:
41+
del sys.modules[_name]
42+
"""
43+
) + textwrap.dedent(body)
44+
return subprocess.run(
45+
[sys.executable, "-c", script], capture_output=True, text=True
46+
)
47+
48+
49+
def test_constants_imports_without_torch():
50+
"""serve.constants must not instantiate TorchTensorSerializer at import time."""
51+
result = _run_without_torch(
52+
"""
53+
from sagemaker.serve.constants import (
54+
Framework,
55+
DEFAULT_SERIALIZERS_BY_FRAMEWORK,
56+
)
57+
58+
assert Framework.PYTORCH in DEFAULT_SERIALIZERS_BY_FRAMEWORK
59+
print("OK")
60+
"""
61+
)
62+
assert "OK" in result.stdout, result.stderr
63+
64+
65+
def test_model_builder_imports_without_torch():
66+
"""ModelBuilder must be importable for API-only use without torch installed."""
67+
result = _run_without_torch(
68+
"""
69+
from sagemaker.serve import ModelBuilder
70+
71+
assert ModelBuilder is not None
72+
print("OK")
73+
"""
74+
)
75+
assert "OK" in result.stdout, result.stderr
76+
77+
78+
def test_pytorch_serializer_still_requires_torch():
79+
"""Resolving the PyTorch entry must still raise a clear error without torch."""
80+
result = _run_without_torch(
81+
"""
82+
from sagemaker.serve.constants import (
83+
Framework,
84+
DEFAULT_SERIALIZERS_BY_FRAMEWORK,
85+
)
86+
87+
serializer_cls, _ = DEFAULT_SERIALIZERS_BY_FRAMEWORK[Framework.PYTORCH]
88+
try:
89+
serializer_cls()
90+
raise AssertionError("expected ImportError")
91+
except ImportError as e:
92+
assert "torch" in str(e)
93+
print("OK")
94+
"""
95+
)
96+
assert "OK" in result.stdout, result.stderr

0 commit comments

Comments
 (0)