Функция jit.load
Функция jit.load загружает сохранённую TorchScript модель из файла.
Она восстанавливает модель в том же состоянии, в котором она была сохранена,
включая все параметры, архитектуру и буферы.
Первым параметром функция принимает путь к файлу или файловый объект.
Вторым параметром можно передать устройство (device), на которое будет загружена модель.
Синтаксис
torch.jit.load(path, [map_location])
Пример
Давайте загрузим ранее сохранённую модель из файла model.pt:
import torch
# Load the saved model
model = torch.jit.load('model.pt')
print(type(model))
Результат выполнения кода:
<class 'torch.jit.RecursiveScriptModule'>
Пример
Давайте загрузим модель на устройство cuda с помощью параметра map_location:
import torch
# Load model to a specific device
model = torch.jit.load('model.pt', map_location='cuda:0')
# Move to CPU if needed
model = torch.jit.load('model.pt', map_location=torch.device('cpu'))
Пример
Давайте загрузим модель и применим её к данным:
import torch
# Load the model
model = torch.jit.load('model.pt')
model.eval()
# Create input tensor
t = torch.tensor([1., 2., 3., 4., 5.])
# Make prediction
res = model(t)
print(res)
Результат выполнения кода:
tensor([1., 2., 3., 4., 5.])
Смотрите также
-
функцию
jit.save,
которая сохраняет TorchScript модель в файл -
функцию
jit.trace,
которая создаёт TorchScript модель путём трассировки -
функцию
jit.script,
которая компилирует функцию или класс в TorchScript -
функцию
hub.load,
которая загружает модель из репозитория PyTorch Hub