Spec-Zone.ru › PyTorch 1

torch.load

torch.load(f, map_location=None, pickle_module=pickle, *, weights_only=False, **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().

Параметры:
  • 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) – указывает, нужно ли ограничить распаковщик загрузкой только тензоров, примитивных типов и словарей
  • pickle_load_args (Any) – (только Python 3) необязательные ключевые аргументы, передаваемые pickle_module.load() и pickle_module.Unpickler(), например, errors=....
Тип возвращаемого значения:

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/1.13/generated/torch.load.html

Spec-Zone.ru

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