Функция load
Функция load используется для загрузки объектов PyTorch из файла, который был сохранён с помощью функции save. Первым параметром функция принимает путь к файлу. Вторым параметром можно передать устройство загрузки (CPU или GPU). Функция возвращает загруженный объект (тензор, модель, словарь и так далее). По умолчанию загрузка происходит на CPU, даже если объект был сохранён на GPU.
Синтаксис
torch.load(file_path, [map_location])
Пример
Давайте сохраним тензор в файл, а затем загрузим его обратно:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
torch.save(t, 'tensor.pt')
loaded_t = torch.load('tensor.pt')
print(loaded_t)
Результат выполнения кода:
tensor([1, 2, 3, 4, 5])
Пример
Загрузим сохранённую модель нейронной сети. Сначала сохраним модель, а потом загрузим её в режиме оценки:
import torch
import torch.nn as nn
model = nn.Linear(10, 5)
torch.save(model.state_dict(), 'model.pt')
loaded_model = nn.Linear(10, 5)
loaded_model.load_state_dict(torch.load('model.pt'))
loaded_model.eval()
print("Model loaded successfully")
Результат выполнения кода:
"Model loaded successfully"
Пример
Загрузим объект на конкретное устройство с помощью параметра map_location. В этом примере мы загружаем тензор на CPU, даже если он был сохранён на GPU:
import torch
# Сохраняем тензор на GPU (если доступен)
if torch.cuda.is_available():
t = torch.tensor([1, 2, 3, 4, 5]).cuda()
torch.save(t, 'gpu_tensor.pt')
# Загружаем на CPU
loaded_t = torch.load('gpu_tensor.pt', map_location=torch.device('cpu'))
print(loaded_t)
Результат выполнения кода:
tensor([1, 2, 3, 4, 5])
Смотрите также
-
функцию
save,
которая сохраняет объекты PyTorch в файл -
функцию
jit.load,
которая загружает скриптовые или трассированные модули -
функцию
hub.load,
которая загружает предобученные модели из репозиториев -
функцию
export.load,
которая загружает экспортированные программы