torch.load
-
torch.load(f, map_location=None, pickle_module=pickle, *, weights_only=True, mmap=None, **pickle_load_args)[исходный код] -
Загружает из файла объект, сохранённый с помощью
torch.save().Предупреждение
torch.load()использует под капотом распаковщик pickle. Никогда не загружайте данные из ненадёжного источника.Подробнее см. в разделе безопасность weights_only.
torch.load()использует средства распаковки pickle в Python, но особым образом обрабатывает хранилища, на которых основаны тензоры. Сначала они десериализуются на CPU, а затем перемещаются на устройство, с которого были сохранены. Если это не удаётся (например, если в среде выполнения отсутствуют определённые устройства), возникает исключение. Однако хранилища можно динамически переназначить на другой набор устройств с помощью аргументаmap_location.Если
map_locationявляется вызываемым объектом, он будет вызван для каждого сериализованного хранилища с двумя аргументами: хранилищем и местоположением. Аргумент хранилища представляет собой результат первоначальной десериализации хранилища, находящийся на CPU. Каждому сериализованному хранилищу соответствует тег местоположения, указывающий устройство, с которого оно было сохранено; этот тег передаётся вторым аргументом вmap_location. Встроенные теги местоположений:'cpu'для тензоров CPU и'cuda:device_id'(например,'cuda:2') для тензоров CUDA.map_locationдолжен возвращать либоNone, либо хранилище. Еслиmap_locationвозвращает хранилище, оно будет использовано в качестве окончательного десериализованного объекта, уже перемещённого на нужное устройство. В противном случаеtorch.load()вернётся к поведению по умолчанию, как если быmap_locationне был указан.Если
map_location— это объектtorch.deviceили строка с тегом устройства, он указывает местоположение, куда следует загрузить все тензоры.В противном случае, если
map_location— это dict, он будет использоваться для переназначения тегов местоположений, встречающихся в файле (ключи), на теги, указывающие, куда поместить хранилища (значения).Пользовательские расширения могут регистрировать собственные теги местоположений и методы присвоения тегов и десериализации с помощью
torch.serialization.register_package().Дополнительные инструменты для работы с контрольной точкой см. в разделе Управление компоновкой.
- Параметры:
-
-
f (str | PathLike[str] | IO[bytes]) – файловый объект (должен реализовывать
read(),readline(),tell()иseek()) либо строка или объект os.PathLike, содержащий имя файла -
map_location (Callable[[Storage, str], Storage] | device | str | dict[str, str] | None) – функция,
torch.device, строка или dict, задающий способ переназначения местоположений хранилищ -
pickle_module (Any | None) – модуль, используемый для распаковки метаданных и объектов (должен совпадать с
pickle_module, использованным для сериализации файла) -
weights_only (bool | None) – указывает, следует ли ограничить распаковщик загрузкой только тензоров, примитивных типов, словарей и типов, добавленных с помощью
torch.serialization.add_safe_globals(). Подробнее см. в разделе torch.load с weights_only=True. Еслиweights_only=True, а контрольная точка содержит разреженные тензоры, всегда проверяются их инварианты (например, границы индексов), чтобы некорректные индексы не привели к чтению за пределами допустимой области. Для каждого разреженного тензора выполняется сканирование сложности O(nnz), которое может быть медленным для больших контрольных точек. -
mmap (bool | None) – указывает, следует ли использовать отображение файла в память вместо загрузки всех хранилищ в память. Обычно хранилища тензоров в файле сначала перемещаются с диска в память CPU, а затем — в местоположение, указанное для них при сохранении или заданное через
map_location. Этот второй шаг не выполняется, если конечное местоположение — CPU. Если установлен флагmmap, на первом шаге хранилища тензоров не копируются с диска в память CPU, а отображаетсяf; это означает, что хранилища тензоров загружаются по требованию при обращении к их данным. -
pickle_load_args (Any) – необязательные именованные аргументы, передаваемые в
pickle_module.load()иpickle_module.Unpickler(); работает только еслиweights_only=False, напримерerrors=....
-
f (str | PathLike[str] | IO[bytes]) – файловый объект (должен реализовывать
- Тип возвращаемого значения:
Примечание
При вызове
torch.load()для файла, содержащего тензоры GPU, по умолчанию эти тензоры будут загружены на GPU. Чтобы избежать резкого увеличения использования памяти GPU при загрузке контрольной точки модели, можно вызватьtorch.load(.., map_location='cpu'), а затемload_state_dict().Примечание
По умолчанию мы декодируем байтовые строки как
utf-8. Это позволяет избежать распространённой ошибкиUnicodeDecodeError: 'ascii' codec can't decode byte 0x...при загрузке в Python 3 файлов, сохранённых в Python 2. Если это значение по умолчанию вам не подходит, можно использовать дополнительный именованный аргументencoding, чтобы указать, как загружать эти объекты. Например,encoding='latin1'декодирует их в строки с использованием кодировкиlatin1, аencoding='bytes'сохраняет их как массивы байтов, которые впоследствии можно декодировать с помощьюbyte_array.decode(...).Пример
>>> torch.load("tensors.pt", weights_only=True) # Load all tensors onto the CPU >>> torch.load( ... "tensors.pt", ... map_location=torch.device("cpu"), ... weights_only=True, ... ) # Load all tensors onto the CPU, using a function >>> torch.load( ... "tensors.pt", ... map_location=lambda storage, loc: storage, ... weights_only=True, ... ) # Load all tensors onto GPU 1 >>> torch.load( ... "tensors.pt", ... map_location=lambda storage, loc: storage.cuda(1), # type: ignore[attr-defined] ... weights_only=True, ... ) # type: ignore[attr-defined] # Map tensors from GPU 1 to GPU 0 >>> torch.load( ... "tensors.pt", ... map_location={"cuda:1": "cuda:0"}, ... weights_only=True, ... ) # Load tensor from io.BytesIO object # Loading from a buffer setting weights_only=False, warning this can be unsafe >>> with open("tensor.pt", "rb") as f: ... buffer = io.BytesIO(f.read()) >>> torch.load(buffer, weights_only=False) # Load a module with 'ascii' encoding for unpickling # Loading from a module setting weights_only=False, warning this can be unsafe >>> torch.load("module.pt", encoding="ascii", weights_only=False)
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.load.html