Skip to content

Commit 705e3ae

Browse files
Improve error message for weights_only load (#129783)
* Improve error message for weights_only load (#129705) As @vmoens pointed out, the current error message does not make the "either/or" between setting `weights_only=False` and using `add_safe_globals` clear enough, and should print the code for the user to call `add_safe_globals` New formatting looks like such In the case that `add_safe_globals` can be used ```python >>> import torch >>> from torch.testing._internal.two_tensor import TwoTensor >>> torch.save(TwoTensor(torch.randn(2), torch.randn(2)), "two_tensor.pt") >>> torch.load("two_tensor.pt", weights_only=True) Traceback (most recent call last): File "<stdin>", line 1, in <module> File "/data/users/mg1998/pytorch/torch/serialization.py", line 1225, in load raise pickle.UnpicklingError(_get_wo_message(str(e))) from None _pickle.UnpicklingError: Weights only load failed. This file can still be loaded, to do so you have two options (1) Re-running `torch.load` with `weights_only` set to `False` will likely succeed, but it can result in arbitrary code execution. Do it only if you got the file from a trusted source. (2) Alternatively, to load with `weights_only=True` please check the recommended steps in the following error message. WeightsUnpickler error: Unsupported global: GLOBAL torch.testing._internal.two_tensor.TwoTensor was not an allowed global by default. Please use `torch.serialization.add_safe_globals([TwoTensor])` to allowlist this global if you trust this class/function. Check the documentation of torch.load to learn more about types accepted by default with weights_only https://pytorch.org/docs/stable/generated/torch.load.html. ``` For other issues (unsupported bytecode) ```python >>> import torch >>> t = torch.randn(2, 3) >>> torch.save(t, "protocol_5.pt", pickle_protocol=5) >>> torch.load("protocol_5.pt", weights_only=True) /data/users/mg1998/pytorch/torch/_weights_only_unpickler.py:359: UserWarning: Detected pickle protocol 5 in the checkpoint, which was not the default pickle protocol used by `torch.load` (2). The weights_only Unpickler might not support all instructions implemented by this protocol, please file an issue for adding support if you encounter this. warnings.warn( Traceback (most recent call last): File "<stdin>", line 1, in <module> File "/data/users/mg1998/pytorch/torch/serialization.py", line 1225, in load raise pickle.UnpicklingError(_get_wo_message(str(e))) from None _pickle.UnpicklingError: Weights only load failed. Re-running `torch.load` with `weights_only` set to `False` will likely succeed, but it can result in arbitrary code execution. Do it only if you got the file from a trusted source. Please file an issue with the following so that we can make `weights_only=True` compatible with your use case: WeightsUnpickler error: Unsupported operand 149 Check the documentation of torch.load to learn more about types accepted by default with weights_only https://pytorch.org/docs/stable/generated/torch.load.html. ``` Old formatting would have been like: ```python Traceback (most recent call last): File "<stdin>", line 1, in <module> File "/data/users/mg1998/pytorch/torch/serialization.py", line 1203, in load raise pickle.UnpicklingError(UNSAFE_MESSAGE + str(e)) from None _pickle.UnpicklingError: Weights only load failed. Re-running `torch.load` with `weights_only` set to `False` will likely succeed, but it can result in arbitrary code execution. Do it only if you get the file from a trusted source. Alternatively, to load with `weights_only` please check the recommended steps in the following error message. WeightsUnpickler error: Unsupported global: GLOBAL torch.testing._internal.two_tensor.TwoTensor was not an allowed global by default. Please use `torch.serialization.add_safe_globals` to allowlist this global if you trust this class/function. ``` Pull Request resolved: #129705 Approved by: https://github.com/albanD, https://github.com/vmoens ghstack dependencies: #129239, #129396, #129509 (cherry picked from commit 45f3e20) * Fix pickle import when rebase onto release/2.4 * Update torch/serialization.py fix bad rebase again --------- Co-authored-by: Mikayla Gawarecki <mikaylagawarecki@gmail.com>
1 parent b26cde4 commit 705e3ae

3 files changed

Lines changed: 57 additions & 15 deletions

File tree

‎test/test_serialization.py‎

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1112,6 +1112,22 @@ def fake_set_state(obj, *args):
11121112
torch.serialization.clear_safe_globals()
11131113
ClassThatUsesBuildInstruction.__setstate__ = None
11141114

1115+
@parametrize("unsafe_global", [True, False])
1116+
def test_weights_only_error(self, unsafe_global):
1117+
sd = {'t': TwoTensor(torch.randn(2), torch.randn(2))}
1118+
pickle_protocol = torch.serialization.DEFAULT_PROTOCOL if unsafe_global else 5
1119+
with BytesIOContext() as f:
1120+
torch.save(sd, f, pickle_protocol=pickle_protocol)
1121+
f.seek(0)
1122+
if unsafe_global:
1123+
with self.assertRaisesRegex(pickle.UnpicklingError,
1124+
r"use `torch.serialization.add_safe_globals\(\[TwoTensor\]\)` to allowlist"):
1125+
torch.load(f, weights_only=True)
1126+
else:
1127+
with self.assertRaisesRegex(pickle.UnpicklingError,
1128+
"file an issue with the following so that we can make `weights_only=True`"):
1129+
torch.load(f, weights_only=True)
1130+
11151131
@parametrize('weights_only', (False, True))
11161132
def test_serialization_math_bits(self, weights_only):
11171133
t = torch.randn(1, dtype=torch.cfloat)

‎torch/_weights_only_unpickler.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -210,8 +210,8 @@ def load(self):
210210
else:
211211
raise RuntimeError(
212212
f"Unsupported global: GLOBAL {full_path} was not an allowed global by default. "
213-
"Please use `torch.serialization.add_safe_globals` to allowlist this global "
214-
"if you trust this class/function."
213+
f"Please use `torch.serialization.add_safe_globals([{name}])` to allowlist "
214+
"this global if you trust this class/function."
215215
)
216216
elif key[0] == NEWOBJ[0]:
217217
args = self.stack.pop()

‎torch/serialization.py‎

Lines changed: 39 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
import functools
44
import os
55
import io
6+
import re
67
import shutil
78
import struct
89
import sys
@@ -994,12 +995,33 @@ def load(
994995
"""
995996
torch._C._log_api_usage_once("torch.load")
996997
UNSAFE_MESSAGE = (
997-
"Weights only load failed. Re-running `torch.load` with `weights_only` set to `False`"
998-
" will likely succeed, but it can result in arbitrary code execution."
999-
" Do it only if you get the file from a trusted source. Alternatively, to load"
1000-
" with `weights_only` please check the recommended steps in the following error message."
1001-
" WeightsUnpickler error: "
998+
"Re-running `torch.load` with `weights_only` set to `False` will likely succeed, "
999+
"but it can result in arbitrary code execution. Do it only if you got the file from a "
1000+
"trusted source."
10021001
)
1002+
DOCS_MESSAGE = (
1003+
"\n\nCheck the documentation of torch.load to learn more about types accepted by default with "
1004+
"weights_only https://pytorch.org/docs/stable/generated/torch.load.html."
1005+
)
1006+
1007+
def _get_wo_message(message: str) -> str:
1008+
pattern = r"GLOBAL (\S+) was not an allowed global by default."
1009+
has_unsafe_global = re.search(pattern, message) is not None
1010+
if has_unsafe_global:
1011+
updated_message = (
1012+
"Weights only load failed. This file can still be loaded, to do so you have two options "
1013+
f"\n\t(1) {UNSAFE_MESSAGE}\n\t(2) Alternatively, to load with `weights_only=True` please check "
1014+
"the recommended steps in the following error message.\n\tWeightsUnpickler error: "
1015+
+ message
1016+
)
1017+
else:
1018+
updated_message = (
1019+
f"Weights only load failed. {UNSAFE_MESSAGE}\n Please file an issue with the following "
1020+
"so that we can make `weights_only=True` compatible with your use case: WeightsUnpickler "
1021+
"error: " + message
1022+
)
1023+
return updated_message + DOCS_MESSAGE
1024+
10031025
if weights_only is None:
10041026
weights_only, warn_weights_only = False, True
10051027
else:
@@ -1071,12 +1093,14 @@ def load(
10711093
overall_storage=overall_storage,
10721094
**pickle_load_args)
10731095
except RuntimeError as e:
1074-
raise pickle.UnpicklingError(UNSAFE_MESSAGE + str(e)) from None
1075-
return _load(opened_zipfile,
1076-
map_location,
1077-
pickle_module,
1078-
overall_storage=overall_storage,
1079-
**pickle_load_args)
1096+
raise pickle.UnpicklingError(_get_wo_message(str(e))) from None
1097+
return _load(
1098+
opened_zipfile,
1099+
map_location,
1100+
pickle_module,
1101+
overall_storage=overall_storage,
1102+
**pickle_load_args,
1103+
)
10801104
if mmap:
10811105
f_name = "" if not isinstance(f, str) else f"{f}, "
10821106
raise RuntimeError("mmap can only be used with files saved with "
@@ -1086,8 +1110,10 @@ def load(
10861110
try:
10871111
return _legacy_load(opened_file, map_location, _weights_only_unpickler, **pickle_load_args)
10881112
except RuntimeError as e:
1089-
raise pickle.UnpicklingError(UNSAFE_MESSAGE + str(e)) from None
1090-
return _legacy_load(opened_file, map_location, pickle_module, **pickle_load_args)
1113+
raise pickle.UnpicklingError(_get_wo_message(str(e))) from None
1114+
return _legacy_load(
1115+
opened_file, map_location, pickle_module, **pickle_load_args
1116+
)
10911117

10921118

10931119
# Register pickling support for layout instances such as

0 commit comments

Comments
 (0)