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

Метод graph

Метод graph класса ScriptModule возвращает граф вычислений модели в низкоуровневом представлении torch._C.Graph. Этот граф содержит все операции, которые будут выполнены при вызове модуля. Метод не принимает параметров и возвращает объект графа, который можно использовать для анализа, оптимизации или визуализации структуры модели.

Синтаксис

graph = script_module.graph

Где script_module - экземпляр класса ScriptModule, полученный после трассировки или компиляции модели.

Пример

Давайте создадим простую модель и получим её граф вычислений:

import torch import torch.nn as nn class SimpleModel(nn.Module): def __init__(self): super(SimpleModel, self).__init__() self.linear = nn.Linear(10, 5) self.relu = nn.ReLU() def forward(self, x): x = self.linear(x) x = self.relu(x) return x model = SimpleModel() example_input = torch.randn(1, 10) scripted_model = torch.jit.trace(model, example_input) graph = scripted_model.graph print(graph)

Результат выполнения кода:

graph(%x.1 : Tensor): %linear_weight : Tensor = prim::Constant[value=<Tensor>]() %linear_bias : Tensor = prim::Constant[value=<Tensor>]() %6 : Tensor = aten::addmm(%linear_bias, %x.1, %linear_weight, %4, %5) %7 : Tensor = aten::relu(%6) return (%7)

Пример

Получим подробную информацию о графе, используя вспомогательные методы:

import torch import torch.nn as nn class MyModel(nn.Module): def __init__(self): super(MyModel, self).__init__() self.fc1 = nn.Linear(5, 3) self.fc2 = nn.Linear(3, 2) def forward(self, x): x = torch.sigmoid(self.fc1(x)) x = self.fc2(x) return x model = MyModel() x = torch.randn(2, 5) scripted = torch.jit.trace(model, x) graph = scripted.graph print("Number of nodes:", len(graph.nodes())) print("Number of inputs:", len(graph.inputs())) print("Number of outputs:", len(graph.outputs())) print("\nAll nodes:") for node in graph.nodes(): print(" -", node.kind())

Результат выполнения кода:

"Number of nodes: 6" "Number of inputs: 1" "Number of outputs: 1" "\nAll nodes:" " - prim::Constant" " - prim::Constant" " - aten::addmm" " - aten::sigmoid" " - aten::addmm" " - prim::Constant"

Пример

Используем граф для модификации модели и её оптимизации. В этом примере мы удалим операцию ReLU из графа:

import torch import torch.nn as nn class ModelWithRelu(nn.Module): def __init__(self): super(ModelWithRelu, self).__init__() self.fc = nn.Linear(4, 4) def forward(self, x): x = self.fc(x) x = torch.relu(x) return x model = ModelWithRelu() example = torch.randn(1, 4) scripted = torch.jit.trace(model, example) print("Original graph:") print(scripted.graph) # Получаем граф и удаляем операцию relu graph = scripted.graph for node in graph.nodes(): if node.kind() == "aten::relu": node.destroy() print("\nModified graph:") print(scripted.graph)

Результат выполнения кода:

"Original graph:" "graph(%x.1 : Tensor):" " %linear_weight : Tensor = prim::Constant[value=<Tensor>]()" " %linear_bias : Tensor = prim::Constant[value=<Tensor>]()" " %5 : Tensor = aten::addmm(%linear_bias, %x.1, %linear_weight, %3, %4)" " %6 : Tensor = aten::relu(%5)" " return (%6)" "\nModified graph:" "graph(%x.1 : Tensor):" " %linear_weight : Tensor = prim::Constant[value=<Tensor>]()" " %linear_bias : Tensor = prim::Constant[value=<Tensor>]()" " %5 : Tensor = aten::addmm(%linear_bias, %x.1, %linear_weight, %3, %4)" " return (%5)"

Смотрите также

  • класс ScriptModule,
    который представляет скомпилированную версию модуля PyTorch
  • метод save,
    который сохраняет ScriptModule в файл
  • метод code,
    который возвращает Python-представление кода модуля
  • функцию graph,
    которая возвращает низкоуровневый граф вычислений модели
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить