РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
649 of 769 menu

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