Функция jit.trace
Функция jit.trace выполняет трассировку модели PyTorch, преобразуя её в скриптовый модуль TorchScript.
Первый параметр функции принимает модель или функцию, которую нужно трассировать.
Второй параметр принимает пример входных данных example_inputs - кортеж тензоров.
Также можно передать необязательные параметры: check_trace для проверки трассировки и strict для контроля строгого режима.
Синтаксис
torch.jit.trace(model, example_inputs, [check_trace], [strict])
Пример
Давайте выполним трассировку простой модели линейного слоя:
import torch
import torch.nn as nn
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.linear = nn.Linear(5, 3)
def forward(self, x):
return self.linear(x)
model = SimpleModel()
example_input = torch.randn(1, 5)
traced_model = torch.jit.trace(model, example_input)
print(traced_model)
Результат выполнения кода:
"TracedModule<...>"
Пример
Трассировка функции с несколькими входными тензорами:
import torch
def add_and_multiply(x, y):
return x + y, x * y
x = torch.tensor([1, 2, 3])
y = torch.tensor([4, 5, 6])
traced_fn = torch.jit.trace(add_and_multiply, (x, y))
res1, res2 = traced_fn(x, y)
print(res1)
print(res2)
Результат выполнения кода:
tensor([5, 7, 9])
tensor([4, 10, 18])
Пример
Использование параметра check_trace для проверки корректности трассировки:
import torch
import torch.nn as nn
class TestModel(nn.Module):
def forward(self, x):
return x.sum(dim=1, keepdim=True)
model = TestModel()
example_input = torch.randn(2, 3)
traced_model = torch.jit.trace(
model,
example_input,
check_trace=True
)
print(traced_model)
Смотрите также
-
функцию
jit.script,
которая выполняет компиляцию модели в TorchScript -
функцию
jit.save,
которая сохраняет скриптовый модуль на диск -
функцию
jit.load,
которая загружает скриптовый модуль с диска -
функцию
jit.trace_module,
которая выполняет трассировку конкретных методов модуля