Класс RecursiveScriptModule
Класс RecursiveScriptModule представляет собой скомпилированную
версию модуля PyTorch, полученную с помощью TorchScript. Экземпляры этого
класса создаются автоматически при использовании функций torch.jit.script
или torch.jit.trace. Такой модуль можно сохранять на диск, загружать
обратно и выполнять без доступа к исходному Python-коду. Основные методы
класса: save для сохранения на диск, load для загрузки,
а также прямой вызов объекта для выполнения прямого прохода.
Синтаксис
# Создание RecursiveScriptModule из существующего модуля
script_module = torch.jit.script(model)
# Сохранение на диск
script_module.save("model.pt")
# Загрузка с диска
loaded = torch.jit.load("model.pt")
# Выполнение прямого прохода
output = script_module(input_tensor)
Пример
Давайте создадим простую модель, скомпилируем её и выполним прямой проход:
import torch
# Создаём простую линейную модель
model = torch.nn.Linear(5, 3)
model.eval()
# Компилируем в TorchScript
script_module = torch.jit.script(model)
# Создаём входной тензор
x = torch.randn(2, 5)
# Выполняем прямой проход
output = script_module(x)
print(output)
Результат выполнения кода:
tensor([
[-0.0895, -0.5332, -0.8045],
[-0.5116, -0.6450, -0.0992],
], grad_fn=<AddmmBackward0>)
Пример
Рассмотрим сохранение скомпилированной модели на диск и её последующую загрузку:
import torch
# Создаём и компилируем модель
model = torch.nn.Sequential(
torch.nn.Linear(10, 20),
torch.nn.ReLU(),
torch.nn.Linear(20, 5)
)
model.eval()
script_module = torch.jit.script(model)
# Сохраняем на диск
script_module.save("model.pt")
# Загружаем с диска
loaded_module = torch.jit.load("model.pt")
# Проверяем работу загруженной модели
x = torch.randn(1, 10)
output = loaded_module(x)
print(output.shape)
Результат выполнения кода:
torch.Size([1, 5])
Пример
RecursiveScriptModule поддерживает доступ к вложенным
подмодулям и их параметрам через точечную нотацию:
import torch
class MyModule(torch.nn.Module):
def __init__(self):
super().__init__()
self.fc1 = torch.nn.Linear(10, 20)
self.fc2 = torch.nn.Linear(20, 5)
def forward(self, x):
x = torch.relu(self.fc1(x))
return self.fc2(x)
model = MyModule()
model.eval()
script_module = torch.jit.script(model)
# Доступ к параметрам подмодулей
weight = script_module.fc1.weight
bias = script_module.fc2.bias
print(weight.shape)
print(bias.shape)
Результат выполнения кода:
torch.Size([20, 10])
torch.Size([5])
Смотрите также
-
функцию
torch.jit.script,
которая компилирует модуль или функцию в TorchScript -
функцию
torch.jit.trace,
которая компилирует модуль на основе примера входных данных -
функцию
torch.save,
которая сохраняет обычный объект PyTorch в файл -
функцию
torch.jit.load,
которая загружает скомпилированный модуль с диска