Spec-Zone.ru › PyTorch 2

torch.load

torch.load(f, map_location=None, pickle_module=pickle, *, weights_only=False, mmap=None, **pickle_load_args) [source]

Загружает объект, сохраненный с помощью torch.save(), из файла.

torch.load() использует средства распаковки 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 является словарем, он будет использоваться для переназначения тегов местоположения, появляющихся в файле (ключи), на теги, определяющие, куда поместить структуры данных (значения).

Пользовательские расширения могут регистрировать свои собственные теги местоположения и методы маркировки и десериализации с помощью torch.serialization.register_package().

Parameters
  • f (Union[str, PathLike, BinaryIO, IO[bytes]]) – объект, похожий на файл (должен реализовывать read(), readline(), tell(), и seek()), или строка или объект os.PathLike, содержащий имя файла
  • map_location (Optional[Union[Callable[[Tensor, str], Tensor], device, str, Dict[str, str]]]) – функция, torch.device, строка или словарь, определяющие, как переназначить местоположения структур данных
  • pickle_module (Optional[Any]) – модуль, используемый для распаковки метаданных и объектов (должен соответствовать pickle_module , использованному для сериализации файла)
  • weights_only (bool) – Указывает, должно ли быть ограничено распаковывание загрузкой только тензоров, примитивных типов и словарей
  • mmap (Optional[bool]) – Указывает, должен ли файл быть отображён в памяти с помощью mmap вместо загрузки всех структур данных в память. Обычно структуры данных тензоров в файле сначала перемещаются из диска в оперативную память процессора CPU, после чего перемещаются в расположение, которое было помечено при сохранении, или указано map_location. Эта вторая операция является пустой операцией, если конечное расположение находится на процессоре CPU. Когда флаг mmap установлен, вместо копирования структур данных тензоров с диска в оперативную память процессора CPU на первом шаге, f отображается в памяти с помощью mmap.
  • pickle_load_args (Any) – (Только Python 3) необязательные ключевые аргументы, передаваемые в pickle_module.load() и pickle_module.Unpickler(), например, errors=....
Return type

Any

Предупреждение

torch.load(), если параметр weights_only не установлен в значение True, неявно использует модуль pickle, который известен своей небезопасностью. Возможно создание вредоночных данных pickle, которые выполнят произвольный код во время распаковки. Никогда не загружайте данные, которые могли поступить из ненадежного источника в небезопасном режиме или которые могли быть повреждены. Загружайте только надёжные данные.

Примечание

При вызове torch.load() на файле, содержащем тензоры GPU, эти тензоры будут загружены в GPU по умолчанию. Можно вызвать torch.load(.., map_location='cpu') и затем load_state_dict() для предотвращения скачка использования оперативной памяти GPU при загрузке контрольной точки модели.

Примечание

По умолчанию, мы декодируем строковые значения байт как utf-8. Это сделано для предотвращения распространённой ошибки UnicodeDecodeError: 'ascii' codec can't decode byte 0x... при загрузке файлов, сохранённых Python 2 в Python 3. Если это значение по умолчанию неверно, можно использовать дополнительный ключевой аргумент encoding для указания, как должны загружаться эти объекты, например, encoding='latin1' декодирует их в строки с использованием кодировки latin1, а encoding='bytes' сохраняет их как массивы байтов, которые можно декодировать позже с помощью byte_array.decode(...).

Пример

>>> torch.load('tensors.pt')
# Load all tensors onto the CPU
>>> torch.load('tensors.pt', map_location=torch.device('cpu'))
# Load all tensors onto the CPU, using a function
>>> torch.load('tensors.pt', map_location=lambda storage, loc: storage)
# Load all tensors onto GPU 1
>>> torch.load('tensors.pt', map_location=lambda storage, loc: storage.cuda(1))
# Map tensors from GPU 1 to GPU 0
>>> torch.load('tensors.pt', map_location={'cuda:1': 'cuda:0'})
# Load tensor from io.BytesIO object
>>> with open('tensor.pt', 'rb') as f:
...     buffer = io.BytesIO(f.read())
>>> torch.load(buffer)
# Load a module with 'ascii' encoding for unpickling
>>> torch.load('module.pt', encoding='ascii')

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.load.html

Spec-Zone.ru

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