Класс ScriptModule
Класс ScriptModule является основным контейнером для модулей,
скомпилированных в TorchScript. Он создаётся путём трассировки или
скриптинга обычного модуля nn.Module с помощью функций
torch.jit.trace или torch.jit.script.
Объекты этого класса могут быть сериализованы, оптимизированы
и выполнены вне среды Python.
Синтаксис
# Создание через трассировку
script_module = torch.jit.trace(module, example_input)
# Создание через скриптинг
script_module = torch.jit.script(module)
# Вызов скомпилированного модуля
output = script_module(input_data)
Пример
Создадим простой модуль и скомпилируем его через трассировку:
import torch
class SimpleModule(torch.nn.Module):
def forward(self, x):
return x + 2
module = SimpleModule()
example = torch.tensor([1, 2, 3])
script_module = torch.jit.trace(module, example)
print(script_module)
Результат выполнения кода:
ScriptModule(
(original_name): SimpleModule
)
Пример
Скомпилируем модуль через скриптинг и выполним его:
import torch
class ComplexModule(torch.nn.Module):
def forward(self, x, y):
return x * y + 1
module = ComplexModule()
script_module = torch.jit.script(module)
t = torch.tensor([1, 2, 3, 4])
res = script_module(t, torch.tensor(2))
print(res)
Результат выполнения кода:
tensor([3, 5, 7, 9])
Пример
Сохраним скомпилированный модуль в файл и загрузим обратно:
import torch
class LinearModule(torch.nn.Module):
def __init__(self):
super().__init__()
self.linear = torch.nn.Linear(3, 2)
def forward(self, x):
return self.linear(x)
module = LinearModule()
script_module = torch.jit.script(module)
script_module.save('model.pt')
loaded_module = torch.jit.load('model.pt')
t = torch.tensor([[1.0, 2.0, 3.0]])
res = loaded_module(t)
print(res)
Результат выполнения кода:
tensor([[-0.7050, -0.3370]], grad_fn=<AddmmBackward0>)
Пример
Получим граф вычислений скомпилированного модуля:
import torch
class GraphModule(torch.nn.Module):
def forward(self, x, y):
return x + y * 2
module = GraphModule()
script_module = torch.jit.script(module)
graph = script_module.graph
print(graph)
Результат выполнения кода:
"graph(%x.1 : Tensor, %y.1 : Tensor):
%2 : Tensor = aten::mul(%y.1, %2)
%3 : Tensor = aten::add(%x.1, %2)
return (%3)"
Смотрите также
-
класс
ScriptModule,
который представляет модуль в формате TorchScript -
метод
save,
который сохраняет скомпилированный модуль в файл -
метод
graph,
который возвращает граф вычислений модуля -
метод
code,
который возвращает исходный код скомпилированного модуля