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=....
-
f (Union[str, PathLike, BinaryIO, IO[bytes]]) – объект, похожий на файл (должен реализовывать
- Return type
Предупреждение
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