Следите за новинками
в нашем Telegram канале. Жми, чтобы подписаться:)
631 of 769 menu
◀ ▶

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 из файла
← →
↑
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить