РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
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 для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить