Метод 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,
которая возвращает низкоуровневый граф вычислений модели