Spec-Zone.ru › PyTorch 2.14

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=....
Тип возвращаемого значения:

Any

Примечание

При вызове 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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API