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=....
-
f (Union[str, PathLike, BinaryIO, IO[bytes]]) – объект типа «поток данных» (должен реализовывать
- Тип возвращаемого значения:
Предупреждение
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