Функция export.load
Функция export.load загружает программу, ранее экспортированную с помощью export.export или export.save.
Она предназначена для восстановления модели в формате ExportedProgram, который включает в себя граф вычислений,
параметры и буферы модели. Первым параметром функция принимает путь к файлу.
Вторым параметром можно указать устройства (например, 'cpu' или 'cuda') для загрузки.
Синтаксис
torch.export.load(file_path, device=None)
Пример
Давайте сохраним модель с помощью export.save и затем загрузим её с помощью export.load:
import torch
from torch.export import export, save, load
class MyModule(torch.nn.Module):
def forward(self, x):
return x + 1
model = MyModule()
example_inputs = (torch.randn(3, 4),)
ep = export(model, example_inputs)
save(ep, "model.pt")
loaded_ep = load("model.pt")
print(loaded_ep.module()(torch.ones(3, 4)))
Результат выполнения кода:
tensor([
[2., 2., 2., 2.],
[2., 2., 2., 2.],
[2., 2., 2., 2.],
])
Пример
Загрузка экспортированной программы на устройство CUDA (если доступно):
import torch
from torch.export import export, save, load
class MyModule(torch.nn.Module):
def forward(self, x):
return x * 2
model = MyModule()
example_inputs = (torch.randn(2, 3),)
ep = export(model, example_inputs)
save(ep, "model_cuda.pt")
device = "cuda" if torch.cuda.is_available() else "cpu"
loaded_ep = load("model_cuda.pt", device=device)
print(loaded_ep.module()(torch.ones(2, 3, device=device)))
Результат выполнения кода (на CPU):
tensor([
[2., 2., 2.],
[2., 2., 2.],
])
Пример
Загрузка модели с сохранёнными параметрами и применение к новым данным:
import torch
from torch.export import export, save, load
class LinearModule(torch.nn.Module):
def __init__(self):
super().__init__()
self.linear = torch.nn.Linear(5, 3)
def forward(self, x):
return self.linear(x)
model = LinearModule()
example_inputs = (torch.randn(2, 5),)
ep = export(model, example_inputs)
save(ep, "linear_model.pt")
loaded_ep = load("linear_model.pt")
t = torch.randn(2, 5)
res = loaded_ep.module()(t)
print(res)
Результат выполнения кода:
tensor([
[ 0.2679, -0.6827, 0.4435],
[ 0.0570, -0.7008, -0.1617],
], grad_fn=<AddmmBackward0>)
Смотрите также
-
функцию
export.save,
которая сохраняет экспортированную программу в файл -
функцию
export.export,
которая создаёт экспортированную программу из модели -
функцию
load,
которая загружает сохранённое состояние модели (state_dict) -
функцию
save,
которая сохраняет объект модели в файл