jit.trace_module
Функция jit.trace_module создает трассированную
копию переданного модуля. Она возвращает экземпляр
ScriptModule, который ведет себя так же, как
оригинальный модуль, но при этом его выполнение
может быть оптимизировано. Первым параметром
передается модуль для трассировки. Вторым
параметром передается словарь, где ключами
являются имена методов модуля, а значениями -
кортежи с входными аргументами для этих методов.
Синтаксис
torch.jit.trace_module(mod, inputs)
Пример
Давайте создадим простой модуль и выполним его трассировку:
import torch
import torch.nn as nn
class MyModule(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(10, 5)
def forward(self, x):
return self.fc(x)
def predict(self, x):
return torch.softmax(self.fc(x), dim=1)
mod = MyModule()
example_input = torch.randn(3, 10)
traced_mod = torch.jit.trace_module(mod, {
'forward': (example_input,),
'predict': (example_input,)
})
print(traced_mod)
Результат выполнения кода:
"MyModule( original_name=MyModule )"
Пример
Выполним трассированную версию модуля для проверки корректности:
import torch
import torch.nn as nn
class MyModule(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(10, 5)
def forward(self, x):
return self.fc(x)
def predict(self, x):
return torch.softmax(self.fc(x), dim=1)
mod = MyModule()
example_input = torch.randn(3, 10)
traced_mod = torch.jit.trace_module(mod, {
'forward': (example_input,),
'predict': (example_input,)
})
res = traced_mod.forward(example_input)
print(res.shape)
Результат выполнения кода:
torch.Size([3, 5])
Пример
Используем функцию для сохранения трассированного модуля в файл:
import torch
import torch.nn as nn
class MyModule(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(10, 5)
def forward(self, x):
return self.fc(x)
def predict(self, x):
return torch.softmax(self.fc(x), dim=1)
mod = MyModule()
example_input = torch.randn(3, 10)
traced_mod = torch.jit.trace_module(mod, {
'forward': (example_input,),
'predict': (example_input,)
})
torch.jit.save(traced_mod, 'traced_model.pt')
print(torch.jit.load('traced_model.pt'))
Результат выполнения кода:
"RecursiveScriptModule( original_name=MyModule )"
Смотрите также
-
функцию
jit.trace,
которая создает трассированную версию функции -
функцию
jit.script,
которая компилирует модуль в график TorchScript -
функцию
save,
которая сохраняет объект PyTorch в файл -
функцию
load,
которая загружает объект PyTorch из файла