torch.fx
Обзор
FX — это набор инструментов для разработчиков, предназначенных для преобразования nn.Module экземпляров. FX состоит из трёх основных компонентов: символического трассировщика, промежуточного представления и генерации Python-кода. Пример работы этих компонентов:
import torch
# Simple module for demonstration
class MyModule(torch.nn.Module):
def __init__(self):
super().__init__()
self.param = torch.nn.Parameter(torch.rand(3, 4))
self.linear = torch.nn.Linear(4, 5)
def forward(self, x):
return self.linear(x + self.param).clamp(min=0.0, max=1.0)
module = MyModule()
from torch.fx import symbolic_trace
# Symbolic tracing frontend - captures the semantics of the module
symbolic_traced : torch.fx.GraphModule = symbolic_trace(module)
# High-level intermediate representation (IR) - Graph representation
print(symbolic_traced.graph)
"""
graph():
%x : [num_users=1] = placeholder[target=x]
%param : [num_users=1] = get_attr[target=param]
%add : [num_users=1] = call_function[target=operator.add](args = (%x, %param), kwargs = {})
%linear : [num_users=1] = call_module[target=linear](args = (%add,), kwargs = {})
%clamp : [num_users=1] = call_method[target=clamp](args = (%linear,), kwargs = {min: 0.0, max: 1.0})
return clamp
"""
# Code generation - valid Python code
print(symbolic_traced.code)
"""
def forward(self, x):
param = self.param
add = x + param; x = param = None
linear = self.linear(add); add = None
clamp = linear.clamp(min = 0.0, max = 1.0); linear = None
return clamp
"""
Символический трассировщик выполняет «символическое выполнение» Python-кода. Он подаёт фиктивные значения, называемые Прокси, через код. Операции над этими Прокси записываются. Дополнительную информацию о символическом трассировании можно найти в документации symbolic_trace() и Tracer.
Промежуточное представление — это контейнер для операций, записанных во время символического трассирования. Оно состоит из списка узлов, которые представляют входные данные функции, вызов функций, методов или экземпляров torch.nn.Module, и возвращаемые значения. Дополнительную информацию об IR можно найти в документации Graph. IR — это формат, на котором применяются преобразования.
Генерация Python-кода — то, что делает FX инструментом преобразования Python в Python (или Модуль в Модуль). Для каждой IR-структуры Graph мы можем сгенерировать допустимый Python-код, соответствующий семантике Graph. Эта функциональность обернута в GraphModule, который является экземпляром torch.nn.Module, содержащим Graph и forward метод, сгенерированный из Graph.
Вместе эта цепочка компонентов (символическое трассирование -> промежуточное представление -> преобразования -> генерация Python-кода) составляет цепочку преобразования Python в Python с помощью FX. Кроме того, эти компоненты могут использоваться по отдельности. Например, символическое трассирование может использоваться изолированно для захвата формы кода в целях анализа (а не преобразования). Генерация кода может использоваться для программно сгенерированных моделей, например, из файла конфигурации. У FX много применений!
Несколько примеров преобразований можно найти в репозитории примеров.
Написание Преобразований
Что такое преобразование FX? По сути, это функция, которая выглядит так.
import torch
import torch.fx
def transform(m: nn.Module,
tracer_class : type = torch.fx.Tracer) -> torch.nn.Module:
# Step 1: Acquire a Graph representing the code in `m`
# NOTE: torch.fx.symbolic_trace is a wrapper around a call to
# fx.Tracer.trace and constructing a GraphModule. We'll
# split that out in our transform to allow the caller to
# customize tracing behavior.
graph : torch.fx.Graph = tracer_class().trace(m)
# Step 2: Modify this Graph or create a new one
graph = ...
# Step 3: Construct a Module to return
return torch.fx.GraphModule(m, graph)
Ваше преобразование будет принимать torch.nn.Module, получать Graph из него, выполнять некоторые изменения и возвращать новый torch.nn.Module. Вы должны рассматривать torch.nn.Module, возвращаемый вашим преобразованием FX, как идентичный обычному torch.nn.Module — вы можете передавать его другому преобразованию FX, TorchScript или запускать его. Гарантия того, что вход и выход вашего преобразования FX — это torch.nn.Module, обеспечит композиционность.
Примечание
Также можно изменить существующий GraphModule вместо создания нового, например:
import torch
import torch.fx
def transform(m : nn.Module) -> nn.Module:
gm : torch.fx.GraphModule = torch.fx.symbolic_trace(m)
# Modify gm.graph
# <...>
# Recompile the forward() method of `gm` from its Graph
gm.recompile()
return gm
Обратите внимание, что вы ДОЛЖНЫ вызвать GraphModule.recompile(), чтобы синхронизировать сгенерированный forward() метод на GraphModule с изменённой Graph.
Учитывая, что вы передали torch.nn.Module, который был прослежен в Graph, у вас есть два основных подхода к созданию нового Graph.
Краткий обзор графов
Полное описание семантики графов можно найти в документации Graph, но мы рассмотрим основы здесь. Graph — это структура данных, которая представляет метод GraphModule. Для этого необходима информация:
- Какие входные данные метода?
- Какие операции выполняются внутри метода?
- Какое значение возвращается (т. е. возвращаемое значение) методом?
Все три этих понятия представлены экземплярами Node. Посмотрим на это на коротком примере:
import torch
import torch.fx
class MyModule(torch.nn.Module):
def __init__(self):
super().__init__()
self.param = torch.nn.Parameter(torch.rand(3, 4))
self.linear = torch.nn.Linear(4, 5)
def forward(self, x):
return torch.topk(torch.sum(
self.linear(x + self.linear.weight).relu(), dim=-1), 3)
m = MyModule()
gm = torch.fx.symbolic_trace(m)
gm.graph.print_tabular()
Здесь мы определяем модуль MyModule для демонстрационных целей, инициализируем его, символично отслеживаем его, а затем вызываем метод Graph.print_tabular() для вывода таблицы, показывающей узлы этого Graph:
opcode | name | target | args | kwargs |
|---|---|---|---|---|
placeholder | x | x | () | {} |
get_attr | linear_weight | linear.weight | () | {} |
call_function | add_1 | <built-in function add> | (x, linear_weight) | {} |
call_module | linear_1 | linear | (add_1,) | {} |
call_method | relu_1 | relu | (linear_1,) | {} |
call_function | sum_1 | <built-in method sum …> | (relu_1,) | {‘dim’: -1} |
call_function | topk_1 | <built-in method topk …> | (sum_1, 3) | {} |
output | output | output | (topk_1,) | {} |
Мы можем использовать эту информацию, чтобы ответить на поставленные выше вопросы.
- Какие входные данные метода? В FX входные данные метода указываются с помощью специальных узлов
placeholder. В данном случае, у нас есть один узелplaceholder, со значениемtargetx, что означает, что у нас есть один аргумент (не self) с именем x. - Какие операции внутри метода? Узлы
get_attr,call_function,call_module, иcall_methodпредставляют операции в методе. Полное описание семантики всех этих узлов можно найти в документацииNode. - Какое возвращаемое значение метода? Возвращаемое значение в
Graphзадаётся специальным узломoutput.
Теперь, зная основы того, как код представлен в FX, мы можем изучить, как мы бы отредактировали Graph.
Обработка графов
Прямая обработка графов
Один из подходов к построению нового Graph — это прямое изменение старого. Для этого мы просто берём Graph, полученный из символического трассирования, и изменяем его. Например, предположим, что мы хотим заменить вызовы torch.add() вызовами torch.mul().
import torch
import torch.fx
# Sample module
class M(torch.nn.Module):
def forward(self, x, y):
return torch.add(x, y)
def transform(m: torch.nn.Module,
tracer_class : type = fx.Tracer) -> torch.nn.Module:
graph : fx.Graph = tracer_class().trace(m)
# FX represents its Graph as an ordered list of
# nodes, so we can iterate through them.
for node in graph.nodes:
# Checks if we're calling a function (i.e:
# torch.add)
if node.op == 'call_function':
# The target attribute is the function
# that call_function calls.
if node.target == torch.add:
node.target = torch.mul
graph.lint() # Does some checks to make sure the
# Graph is well-formed.
return fx.GraphModule(m, graph)
Мы также можем выполнять более сложные переписывания Graph, такие как удаление или добавление узлов. Для этих преобразований FX предоставляет вспомогательные функции для обработки графа, которые можно найти в документации Graph. Пример использования этих API для добавления вызова torch.relu() показан ниже.
# Specifies the insertion point. Any nodes added to the
# Graph within this scope will be inserted after `node`
with traced.graph.inserting_after(node):
# Insert a new `call_function` node calling `torch.relu`
new_node = traced.graph.call_function(
torch.relu, args=(node,))
# We want all places that used the value of `node` to
# now use that value after the `relu` call we've added.
# We use the `replace_all_uses_with` API to do this.
node.replace_all_uses_with(new_node)
Для простых преобразований, состоящих только из подстановок, вы также можете использовать переписыватель подграфов subgraph rewriter.
Переписывание подграфов с помощью replace_pattern()
FX также предоставляет ещё один уровень автоматизации поверх непосредственного управления графом. API replace_pattern() по сути является инструментом «найти/заменить» для редактирования графов Graph. Он позволяет указать функцию pattern и функцию replacement и будет отслеживать эти функции, искать экземпляры группы операций в графе pattern и заменять эти экземпляры копиями графа replacement. Это может помочь значительно автоматизировать рутинную работу с графами, которая может стать неуправляемой по мере усложнения преобразований.
Примеры манипулирования графом
- Замена одной операции
- Слияние Conv/Batch Norm
- replace_pattern: Базовое использование
- Квантование
- Инвертирование преобразования
Прокси/Переотслеживание
Другой способ манипулирования графами Graph — это повторное использование механизма Proxy, используемого при символическом отслеживании. Например, предположим, что мы хотим написать преобразование, которое разбивает функции PyTorch на более мелкие операции. Оно бы преобразовало каждый вызов F.relu(x) в (x > 0) * x. Одним из вариантов было бы выполнить необходимое переписывание графа, чтобы вставить сравнение и умножение после F.relu, а затем очистить исходный F.relu. Однако мы можем автоматизировать этот процесс, используя объекты Proxy, чтобы автоматически записывать операции в граф Graph.
Для использования этого метода мы пишем операции, которые мы хотим вставить, как обычный код PyTorch, и вызываем этот код с объектами Proxy в качестве аргументов. Эти объекты Proxy будут захватывать операции, выполняемые над ними, и добавлять их в граф Graph.
# Note that this decomposition rule can be read as regular Python
def relu_decomposition(x):
return (x > 0) * x
decomposition_rules = {}
decomposition_rules[F.relu] = relu_decomposition
def decompose(model: torch.nn.Module,
tracer_class : type = fx.Tracer) -> torch.nn.Module:
"""
Decompose `model` into smaller constituent operations.
Currently,this only supports decomposing ReLU into its
mathematical definition: (x > 0) * x
"""
graph : fx.Graph = tracer_class().trace(model)
new_graph = fx.Graph()
env = {}
tracer = torch.fx.proxy.GraphAppendingTracer(new_graph)
for node in graph.nodes:
if node.op == 'call_function' and node.target in decomposition_rules:
# By wrapping the arguments with proxies,
# we can dispatch to the appropriate
# decomposition rule and implicitly add it
# to the Graph by symbolically tracing it.
proxy_args = [
fx.Proxy(env[x.name], tracer) if isinstance(x, fx.Node) else x for x in node.args]
output_proxy = decomposition_rules[node.target](*proxy_args)
# Operations on `Proxy` always yield new `Proxy`s, and the
# return value of our decomposition rule is no exception.
# We need to extract the underlying `Node` from the `Proxy`
# to use it in subsequent iterations of this transform.
new_node = output_proxy.node
env[node.name] = new_node
else:
# Default case: we don't have a decomposition rule for this
# node, so just copy the node over into the new graph.
new_node = new_graph.node_copy(node, lambda x: env[x.name])
env[node.name] = new_node
return fx.GraphModule(model, new_graph)
Помимо избежания явного манипулирования графом, использование объектов Proxy также позволяет определять правила переписывания как обычный код Python. Для преобразований, требующих большого количества правил переписывания (таких как vmap или grad), это часто повышает читаемость и поддерживаемость правил. Обратите внимание, что при вызове Proxy мы также передавали трейсер, указывающий на базовую переменную graph. Это делается в случае, если операции в графе являются n-арными (например, add — бинарный оператор), вызов Proxy не создает несколько экземпляров трейсера графа, что может привести к неожиданным ошибкам во время выполнения. Мы рекомендуем этот метод использования объектов Proxy, особенно когда невозможно безопасно предположить, что базовые операторы являются унарными.
Пример практического использования объектов Proxy для манипулирования графами Graph можно найти здесь.
Шаблон интерпретатора
Полезным шаблоном организации кода в FX является перебор всех узлов Node в графе Graph и выполнение их. Это можно использовать для нескольких целей, включая аналитику значений, протекающих через граф во время выполнения, или преобразование кода посредством переотслеживания с использованием объектов Proxy. Например, предположим, что мы хотим выполнить GraphModule и записывать свойства формы и типа данных тензора torch.Tensor для узлов во время выполнения. Это может выглядеть так:
import torch
import torch.fx
from torch.fx.node import Node
from typing import Dict
class ShapeProp:
"""
Shape propagation. This class takes a `GraphModule`.
Then, its `propagate` method executes the `GraphModule`
node-by-node with the given arguments. As each operation
executes, the ShapeProp class stores away the shape and
element type for the output values of each operation on
the `shape` and `dtype` attributes of the operation's
`Node`.
"""
def __init__(self, mod):
self.mod = mod
self.graph = mod.graph
self.modules = dict(self.mod.named_modules())
def propagate(self, *args):
args_iter = iter(args)
env : Dict[str, Node] = {}
def load_arg(a):
return torch.fx.graph.map_arg(a, lambda n: env[n.name])
def fetch_attr(target : str):
target_atoms = target.split('.')
attr_itr = self.mod
for i, atom in enumerate(target_atoms):
if not hasattr(attr_itr, atom):
raise RuntimeError(f"Node referenced nonexistant target {'.'.join(target_atoms[:i])}")
attr_itr = getattr(attr_itr, atom)
return attr_itr
for node in self.graph.nodes:
if node.op == 'placeholder':
result = next(args_iter)
elif node.op == 'get_attr':
result = fetch_attr(node.target)
elif node.op == 'call_function':
result = node.target(*load_arg(node.args), **load_arg(node.kwargs))
elif node.op == 'call_method':
self_obj, *args = load_arg(node.args)
kwargs = load_arg(node.kwargs)
result = getattr(self_obj, node.target)(*args, **kwargs)
elif node.op == 'call_module':
result = self.modules[node.target](*load_arg(node.args), **load_arg(node.kwargs))
# This is the only code specific to shape propagation.
# you can delete this `if` branch and this becomes
# a generic GraphModule interpreter.
if isinstance(result, torch.Tensor):
node.shape = result.shape
node.dtype = result.dtype
env[node.name] = result
return load_arg(self.graph.result)
Как видите, полноценный интерпретатор для FX несложен, но может быть очень полезным. Для упрощения использования этого шаблона мы предоставляем класс Interpreter, который обобщает вышеупомянутую логику таким образом, что определённые аспекты выполнения интерпретатора могут быть переопределены посредством переопределения методов.
Помимо выполнения операций, мы также можем сгенерировать новый Graph, передавая значения Proxy через интерпретатор. Аналогично, мы предоставляем класс Transformer, который обобщает этот шаблон. Transformer ведет себя аналогично Interpreter, но вместо вызова метода run для получения конкретного выходного значения из модуля, вы вызываете метод Transformer.transform() для возврата нового GraphModule, который был подвергнут любым правилам трансформации, которые вы установили как переопределенные методы.
Примеры шаблона интерпретатора
Отладка
Введение
Часто при написании преобразований наш код не совсем правильный. В этом случае нам может потребоваться отладка. Ключ — работать в обратном порядке: сначала проверьте результаты вызова сгенерированного модуля, чтобы подтвердить или опровергнуть правильность. Затем проверьте и отладьте сгенерированный код. Затем отладьте процесс преобразований, которые привели к сгенерированному коду.
Если вы не знакомы с отладчиками, см. дополнительный раздел Доступные отладчики.
Распространённые ошибки при создании преобразований
- Недетерминированный порядок итерирования
set. В Python тип данныхsetнеупорядочен. Использованиеsetдля хранения коллекций объектов, таких какNode, например, может привести к неожиданной недетерминированности. Пример — итерация по наборуNodeдля их вставки вGraph. Поскольку тип данныхsetнеупорядочен, порядок операций в выходной программе будет недетерминированным и может меняться при каждом вызове программы. Рекомендуемым вариантом является использование типа данныхdict, который является упорядоченным по вставке начиная с Python 3.7 (и cPython 3.6). Типdictможно использовать аналогично множеству, храня значения, которые нужно исключить из повторений, в ключахdict.
Проверка корректности модулей
Так как выход большинства глубоких нейронных модулей состоит из тензоров с плавающей точкой torch.Tensor, проверка эквивалентности результатов работы двух torch.nn.Module не так проста, как проверка на точное равенство. Рассмотрим пример:
import torch
import torch.fx
import torchvision.models as models
def transform(m : torch.nn.Module) -> torch.nn.Module:
gm = torch.fx.symbolic_trace(m)
# Imagine we're doing some transforms here
# <...>
gm.recompile()
return gm
resnet18 = models.resnet18()
transformed_resnet18 = transform(resnet18)
input_image = torch.randn(5, 3, 224, 224)
assert resnet18(input_image) == transformed_resnet18(input_image)
"""
RuntimeError: Boolean value of Tensor with more than one value is ambiguous
"""
Здесь мы попытались проверить равенство значений двух глубоких нейронных моделей с оператором == равенства. Однако это некорректно, так как оператор возвращает тензор, а не булево значение, и для сравнения значений с плавающей точкой необходимо использовать погрешность (или ε), чтобы учесть некоммутативность операций с плавающей точкой (подробнее об этом см. здесь). Вместо этого мы можем использовать torch.allclose(), который даст приблизительное сравнение с учётом относительной и абсолютной пороговой погрешности:
assert torch.allclose(resnet18(input_image), transformed_resnet18(input_image))
Это первый инструмент в нашем арсенале для проверки того, что преобразованные модули ведут себя так, как мы ожидаем, по сравнению с эталонной реализацией.
Отладка сгенерированного кода
Так как FX генерирует функцию forward() для модулей GraphModule, использование традиционных методов отладки, таких как операторы print или pdb, не так просто. К счастью, у нас есть несколько методов отладки сгенерированного кода.
Использование pdb
Вызовите pdb для входа в запущенную программу. Хотя код, представляющий Graph, не находится в файле исходного кода, мы все равно можем вручную войти в него с помощью pdb при вызове прохода вперёд.
import torch
import torch.fx
import torchvision.models as models
def my_pass(inp: torch.nn.Module, tracer_class : type = fx.Tracer) -> torch.nn.Module:
graph = tracer_class().trace(inp)
# Transformation logic here
# <...>
# Return new Module
return fx.GraphModule(inp, graph)
my_module = models.resnet18()
my_module_transformed = my_pass(my_module)
input_value = torch.randn(5, 3, 224, 224)
# When this line is executed at runtime, we will be dropped into an
# interactive `pdb` prompt. We can use the `step` or `s` command to
# step into the execution of the next line
import pdb; pdb.set_trace()
my_module_transformed(input_value)
Вывести сгенерированный код
Если вам нужно запускать один и тот же код несколько раз, то проход к нужному коду с помощью pdb может быть довольно утомительным. В этом случае один из подходов — просто скопировать и вставить сгенерированный forward проход в свой код и изучить его оттуда.
# Assume that `traced` is a GraphModule that has undergone some
# number of transforms
# Copy this code for later
print(traced)
# Print the code generated from symbolic tracing. This outputs:
"""
def forward(self, y):
x = self.x
add_1 = x + y; x = y = None
return add_1
"""
# Subclass the original Module
class SubclassM(M):
def __init__(self):
super().__init__()
# Paste the generated `forward` function (the one we printed and
# copied above) here
def forward(self, y):
x = self.x
add_1 = x + y; x = y = None
return add_1
# Create an instance of the original, untraced Module. Then, create an
# instance of the Module with the copied `forward` function. We can
# now compare the output of both the original and the traced version.
pre_trace = M()
post_trace = SubclassM()
Использование функции to_folder из модуля GraphModule
GraphModule.to_folder() — это метод в GraphModule, который позволяет выгрузить сгенерированный код FX в папку. Хотя копирование прохода вперёд в код часто достаточно, как в Вывести сгенерированный код, может быть проще изучить модули и параметры с помощью to_folder.
m = symbolic_trace(M())
m.to_folder("foo", "Bar")
from foo import Bar
y = Bar()
После запуска приведенного выше примера, мы можем посмотреть на код внутри foo/module.py и изменить его по своему желанию (например, добавив print инструкции или используя pdb) для отладки сгенерированного кода.
Отладка преобразования
Теперь, когда мы определили, что преобразование создаёт неверный код, пришло время отладить само преобразование. Сначала мы проверим раздел Ограничения символического трассирования в документации. После того, как мы убедимся, что трассирование работает как ожидается, цель заключается в том, чтобы выяснить, что пошло не так во время нашего GraphModule преобразования. Быстрый ответ может быть в разделе Написание преобразований, но, если нет, существуют несколько способов изучить наш отслеженный модуль:
# Sample Module
class M(torch.nn.Module):
def forward(self, x, y):
return x + y
# Create an instance of `M`
m = M()
# Symbolically trace an instance of `M` (returns a GraphModule). In
# this example, we'll only be discussing how to inspect a
# GraphModule, so we aren't showing any sample transforms for the
# sake of brevity.
traced = symbolic_trace(m)
# Print the code produced by tracing the module.
print(traced)
# The generated `forward` function is:
"""
def forward(self, x, y):
add = x + y; x = y = None
return add
"""
# Print the internal Graph.
print(traced.graph)
# This print-out returns:
"""
graph():
%x : [num_users=1] = placeholder[target=x]
%y : [num_users=1] = placeholder[target=y]
%add : [num_users=1] = call_function[target=operator.add](args = (%x, %y), kwargs = {})
return add
"""
# Print a tabular representation of the internal Graph.
traced.graph.print_tabular()
# This gives us:
"""
opcode name target args kwargs
------------- ------ ----------------------- ------ --------
placeholder x x () {}
placeholder y y () {}
call_function add <built-in function add> (x, y) {}
output output output (add,) {}
"""
Используя описанные выше вспомогательные функции, мы можем сравнить наш отслеженный Модуль до и после применения наших преобразований. Иногда простое визуальное сравнение достаточно, чтобы отследить ошибку. Если всё ещё не ясно, что не так, отладчик, такой как pdb, может стать хорошим следующим шагом.
На основании примера выше, рассмотрим следующий код:
# Sample user-defined function
def transform_graph(module: torch.nn.Module, tracer_class : type = fx.Tracer) -> torch.nn.Module:
# Get the Graph from our traced Module
g = tracer_class().trace(module)
"""
Transformations on `g` go here
"""
return fx.GraphModule(module, g)
# Transform the Graph
transformed = transform_graph(traced)
# Print the new code after our transforms. Check to see if it was
# what we expected
print(transformed)
Используя вышеприведённый пример, предположим, что вызов print(traced) показал нам ошибку в наших преобразованиях. Мы хотим найти причину ошибки с помощью отладчика. Мы начинаем сессию отладки pdb. Мы можем увидеть, что происходит во время преобразования, остановившись на transform_graph(traced), а затем нажав s для «входа» в вызов transform_graph(traced).
Мы также можем добиться успеха, отредактировав метод print_tabular для вывода различных атрибутов узлов в графе. (Например, мы можем захотеть увидеть атрибуты узла input_nodes и users.)
Доступные отладчики
Наиболее распространённый отладчик Python — pdb. Вы можете запустить свою программу в «режиме отладки» с помощью pdb, введя python -m pdb FILENAME.py в командной строке, где FILENAME — имя файла, который вы хотите отладить. После этого вы можете использовать pdb команды отладчика для пошагового перемещения по вашей запущенной программе. Обычно вы устанавливаете точку останова (b LINE-NUMBER) при запуске pdb, а затем вызываете c для запуска программы до этой точки. Это позволит вам избежать прохождения по каждой строке выполнения (с помощью s или n) для достижения желаемой части кода. В качестве альтернативы, вы можете написать import pdb; pdb.set_trace() перед строкой, на которой хотите установить точку останова. Если вы добавите pdb.set_trace(), ваша программа автоматически запустится в режиме отладки при запуске. (Другими словами, вы можете просто ввести python FILENAME.py в командной строке вместо python -m pdb FILENAME.py.) После запуска файла в режиме отладки вы можете переходить по коду и изучать внутреннее состояние вашей программы с помощью определённых команд. В интернете есть множество отличных учебников по pdb , включая учебник RealPython «Python Debugging With Pdb».
IDE, такие как PyCharm или VSCode, обычно имеют встроенный отладчик. В вашей IDE вы можете либо а) использовать pdb, открыв окно терминала в вашей IDE (например, View → Terminal в VSCode), либо б) использовать встроенный отладчик (обычно графический оболочка вокруг pdb).
Ограничения символического трассирования
FX использует систему символического трассирования (также известного как символическое выполнение) для захвата семантики программ в преобразуемой/анализируемой форме. Система является трассированием, так как она выполняет программу (на самом деле torch.nn.Module или функцию) для записи операций. Она является символической, так как данные, протекающие через программу во время этого выполнения, не являются реальными данными, а символами (Proxy в терминологии FX).
Хотя символическое трассирование работает для большинства кода нейронных сетей, оно имеет некоторые ограничения.
Динамическое управление потоком
Основное ограничение символического трассирования заключается в том, что оно в настоящее время не поддерживает динамическое управление потоком. То есть циклы или if операторы, где условие может зависеть от входных значений программы.
Например, рассмотрим следующую программу:
def func_to_trace(x):
if x.sum() > 0:
return torch.relu(x)
else:
return torch.neg(x)
traced = torch.fx.symbolic_trace(func_to_trace)
"""
<...>
File "dyn.py", line 6, in func_to_trace
if x.sum() > 0:
File "pytorch/torch/fx/proxy.py", line 155, in __bool__
return self.tracer.to_bool(self)
File "pytorch/torch/fx/proxy.py", line 85, in to_bool
raise TraceError('symbolically traced variables cannot be used as inputs to control flow')
torch.fx.proxy.TraceError: symbolically traced variables cannot be used as inputs to control flow
"""
Условие if оператора зависит от значения x.sum(), которое зависит от значения x, входной функции. Поскольку x может изменяться (то есть, если вы передаёте новый тензор ввода в отслеживаемую функцию), это динамическое управление потоком. Трассировка возвращается вверх по вашему коду, чтобы показать, где возникает эта ситуация.
Статическое управление потоком
С другой стороны, поддерживается так называемое статическое управление потоком. Статический поток управления — это циклы или if операторы, значение которых не может изменяться при вызовах. Как правило, в программах PyTorch такой поток управления возникает для кода, принимающего решения об архитектуре модели на основе гиперпараметров. В качестве конкретного примера:
import torch
import torch.fx
class MyModule(torch.nn.Module):
def __init__(self, do_activation : bool = False):
super().__init__()
self.do_activation = do_activation
self.linear = torch.nn.Linear(512, 512)
def forward(self, x):
x = self.linear(x)
# This if-statement is so-called static control flow.
# Its condition does not depend on any input values
if self.do_activation:
x = torch.relu(x)
return x
without_activation = MyModule(do_activation=False)
with_activation = MyModule(do_activation=True)
traced_without_activation = torch.fx.symbolic_trace(without_activation)
print(traced_without_activation.code)
"""
def forward(self, x):
linear_1 = self.linear(x); x = None
return linear_1
"""
traced_with_activation = torch.fx.symbolic_trace(with_activation)
print(traced_with_activation.code)
"""
import torch
def forward(self, x):
linear_1 = self.linear(x); x = None
relu_1 = torch.relu(linear_1); linear_1 = None
return relu_1
"""
Оператор if if self.do_activation не зависит от входных данных функции, следовательно, он статический. do_activation можно рассматривать как гиперпараметр, и трассировки различных экземпляров Module с разными значениями этого параметра имеют разный код. Это допустимый шаблон, поддерживаемый символическим трассированием.
Многие случаи динамического управления потоком являются семантически статическим управлением потоком. Эти случаи можно сделать совместимыми с символическим трассированием, удалив зависимости данных от входных значений, например, перенеся значения в Module атрибуты или привязав конкретные значения к аргументам во время символического трассирования:
def f(x, flag):
if flag: return x
else: return x*2
fx.symbolic_trace(f) # Fails!
fx.symbolic_trace(f, concrete_args={'flag': True})
В случае истинного динамического управления потоком, участки программы, содержащие этот код, могут быть отслежены как вызовы метода (см. Настройка трассировки с помощью класса Tracer) или функции (см. wrap()) вместо трассировки через них.
Функции, не являющиеся функциями torch
FX использует __torch_function__ в качестве механизма перехвата вызовов (см. техническое описание для получения дополнительной информации об этом). Некоторые функции, такие как встроенные функции Python или функции в модуле math, не покрываются __torch_function__, но мы все равно хотели бы захватить их в символическом трассировании. Например:
import torch
import torch.fx
from math import sqrt
def normalize(x):
"""
Normalize `x` by the size of the batch dimension
"""
return x / sqrt(len(x))
# It's valid Python code
normalize(torch.rand(3, 4))
traced = torch.fx.symbolic_trace(normalize)
"""
<...>
File "sqrt.py", line 9, in normalize
return x / sqrt(len(x))
File "pytorch/torch/fx/proxy.py", line 161, in __len__
raise RuntimeError("'len' is not supported in symbolic tracing by default. If you want "
RuntimeError: 'len' is not supported in symbolic tracing by default. If you want this call to be recorded, please call torch.fx.wrap('len') at module scope
"""
Ошибка сообщает нам, что встроенная функция len не поддерживается. Мы можем сделать так, чтобы такие функции записывались в трассировку как прямые вызовы, используя API wrap():
torch.fx.wrap('len')
torch.fx.wrap('sqrt')
traced = torch.fx.symbolic_trace(normalize)
print(traced.code)
"""
import math
def forward(self, x):
len_1 = len(x)
sqrt_1 = math.sqrt(len_1); len_1 = None
truediv = x / sqrt_1; x = sqrt_1 = None
return truediv
"""
Настройка трассировки с помощью класса Tracer
Класс Tracer — это класс, лежащий в основе реализации symbolic_trace. Поведение трассировки можно настроить, создав подкласс Tracer, как показано ниже:
class MyCustomTracer(torch.fx.Tracer):
# Inside here you can override various methods
# to customize tracing. See the `Tracer` API
# reference
pass
# Let's use this custom tracer to trace through this module
class MyModule(torch.nn.Module):
def forward(self, x):
return torch.relu(x) + torch.ones(3, 4)
mod = MyModule()
traced_graph = MyCustomTracer().trace(mod)
# trace() returns a Graph. Let's wrap it up in a
# GraphModule to make it runnable
traced = torch.fx.GraphModule(mod, traced_graph)
Листовые модули
Листовые модули — это модули, которые появляются как вызовы в символической трассировке, а не отслеживаются. Набор стандартных torch.nn экземпляров модулей. Например:
class MySpecialSubmodule(torch.nn.Module):
def forward(self, x):
return torch.neg(x)
class MyModule(torch.nn.Module):
def __init__(self):
super().__init__()
self.linear = torch.nn.Linear(3, 4)
self.submod = MySpecialSubmodule()
def forward(self, x):
return self.submod(self.linear(x))
traced = torch.fx.symbolic_trace(MyModule())
print(traced.code)
# `linear` is preserved as a call, yet `submod` is traced though.
# This is because the default set of "Leaf Modules" includes all
# standard `torch.nn` modules.
"""
import torch
def forward(self, x):
linear_1 = self.linear(x); x = None
neg_1 = torch.neg(linear_1); linear_1 = None
return neg_1
"""
Набор листовых модулей можно настроить, переопределив Tracer.is_leaf_module().
Разное
-
Конструкторы тензоров (например,
torch.zeros,torch.ones,torch.rand,torch.randn,torch.sparse_coo_tensor) в настоящее время не отслеживаются.- Детерминированные конструкторы (
zeros,ones) можно использовать, и значение, которое они производят, будет встроено в трассировку как константа. Это проблема только в том случае, если аргументы этих конструкторов ссылаются на динамические размеры входных данных. В этом случаеones_likeилиzeros_likeмогут быть жизнеспособной заменой. - Для недетерминированных конструкторов (
rand,randn) в трассировку будет встроено единственное случайное значение. Скорее всего, это не является желаемым поведением. Одним из решений является обертываниеtorch.randnв функциюtorch.fx.wrapи вызов этой функции вместо этого.
@torch.fx.wrap def torch_randn(x, shape): return torch.randn(shape) def f(x): return x + torch_randn(x, 5) fx.symbolic_trace(f)- Это поведение может быть исправлено в будущих версиях.
- Детерминированные конструкторы (
-
Аннотации типов
- Аннотации типов в стиле Python 3 (например,
func(x : torch.Tensor, y : int) -> torch.Tensor) поддерживаются и будут сохранены символической трассировкой. - Аннотации типов в стиле комментариев Python 2
# type: (torch.Tensor, int) -> torch.Tensorв настоящее время не поддерживаются. - Аннотации локальных имен внутри функции в настоящее время не поддерживаются.
- Аннотации типов в стиле Python 3 (например,
-
Особенности использования флага
trainingи подмодулей- При использовании функционалов, таких как
torch.nn.functional.dropout, часто аргумент training передается какself.training. Во время трассировки FX это, скорее всего, будет встроено как константа.
import torch import torch.fx class DropoutRepro(torch.nn.Module): def forward(self, x): return torch.nn.functional.dropout(x, training=self.training) traced = torch.fx.symbolic_trace(DropoutRepro()) print(traced.code) """ def forward(self, x): dropout = torch.nn.functional.dropout(x, p = 0.5, training = True, inplace = False); x = None return dropout """ traced.eval() x = torch.randn(5, 3) torch.testing.assert_close(traced(x), x) """ AssertionError: Tensor-likes are not close! Mismatched elements: 15 / 15 (100.0%) Greatest absolute difference: 1.6207983493804932 at index (0, 2) (up to 1e-05 allowed) Greatest relative difference: 1.0 at index (0, 0) (up to 0.0001 allowed) """- Однако, при использовании стандартного подмодуля
nn.Dropout(), флаг training инкапсулирован и — из-за сохранения модели объектаnn.Module— может быть изменен.
class DropoutRepro2(torch.nn.Module): def __init__(self): super().__init__() self.drop = torch.nn.Dropout() def forward(self, x): return self.drop(x) traced = torch.fx.symbolic_trace(DropoutRepro2()) print(traced.code) """ def forward(self, x): drop = self.drop(x); x = None return drop """ traced.eval() x = torch.randn(5, 3) torch.testing.assert_close(traced(x), x) - При использовании функционалов, таких как
- Из-за этого различия рекомендуется отмечать модули, которые динамически взаимодействуют с флагом
training, как листовые модули.
Справочник API
-
torch.fx.symbolic_trace(root, concrete_args=None)[source] -
API символической трассировки
Принимая на вход модуль
nn.Moduleили экземпляр функцииroot, эта функция вернет модульGraphModule, созданный путем записи операций, наблюдаемых при трассировке черезroot.concrete_argsпозволяет частично специализировать функцию, независимо от того, удаляются ли управляющие потоки или структуры данных.Например:
def f(a, b): if b == True: return a else: return a*2FX обычно не может отследить это из-за наличия управляющих потоков. Однако мы можем использовать
concrete_argsдля специализации по значениюbдля отслеживания этого:f = fx.symbolic_trace(f, concrete_args={'b': False}) assert f(3, False) == 6Обратите внимание, что, хотя вы по-прежнему можете передавать разные значения
b, они будут проигнорированы.Мы также можем использовать
concrete_argsдля исключения обработки структур данных из нашей функции. Это использует pytrees для сглаживания входных данных. Чтобы избежать чрезмерной специализации, передайтеfx.PHдля значений, которые не следует специализировать. Например:def f(x): out = 0 for v in x.values(): out += v return out f = fx.symbolic_trace(f, concrete_args={'x': {'a': fx.PH, 'b': fx.PH, 'c': fx.PH}}) assert f({'a': 1, 'b': 2, 'c': 4}) == 7- Параметры
-
- root (Union[torch.nn.Module, Callable]) – Модуль или функция, которые нужно отследить и преобразовать в представление графа.
- concrete_args (Optional[Dict[str, any]]) – Входные данные для частичной специализации
- Возвращает
-
Модуль, созданный из записанных операций из
root. - Тип возвращаемого значения
Примечание
Обратная совместимость для этого API гарантирована.
-
torch.fx.wrap(fn_or_name)[source] -
Эта функция может вызываться на уровне модуля, чтобы зарегистрировать fn_or_name как «листовую функцию». «Листовая функция» будет сохранена как узел CallFunction в трассировке FX вместо того, чтобы быть прослеженной:
# foo/bar/baz.py def my_custom_function(x, y): return x * x + y * y torch.fx.wrap('my_custom_function') def fn_to_be_traced(x, y): # When symbolic tracing, the below call to my_custom_function will be inserted into # the graph rather than tracing it. return my_custom_function(x, y)Эта функция также может быть эквивалентно использована как декоратор:
# foo/bar/baz.py @torch.fx.wrap def my_custom_function(x, y): return x * x + y * yОборачиваемая функция может рассматриваться как «листовая функция», аналогично понятию «листовых модулей», то есть это функции, которые остаются вызовами в трассировке FX, а не прослеживаются.
- Параметры
-
fn_or_name (Union[str, Callable]) – Функция или имя глобальной функции, которую нужно вставить в граф при ее вызове
Примечание
Обратная совместимость для этого API гарантирована.
-
class torch.fx.GraphModule(*args, **kwargs)[source] -
GraphModule — это модуль nn.Module, сгенерированный из fx.Graph. У GraphModule есть атрибут
graph, а также атрибутыcodeиforward, сгенерированные из этогоgraph.Предупреждение
При повторной присвоении
graph,codeиforwardбудут автоматически перегенерированы. Однако, если вы редактируете содержимоеgraphбез повторного присвоения самого атрибутаgraph, необходимо вызватьrecompile()для обновления сгенерированного кода.Примечание
Обратная совместимость для этого API гарантирована.
-
__init__(root, graph, class_name='GraphModule')[source] -
Создаёт GraphModule.
- Параметры
-
-
root (Union[torch.nn.Module, Dict[str, Any]) –
rootможет быть экземпляром nn.Module или словарем, сопоставляющим строки с любым типом атрибута. В случае, еслиrootявляется модулем, любые ссылки на объекты, основанные на модулях (через полное имя), в полеtargetузлов графа будут скопированы из соответствующего места в иерархии модулейrootв иерархию модулей GraphModule. В случае, еслиrootявляется словарем, полное имя, найденное в полеtargetузла, будет найдено непосредственно в ключах словаря. Объект, сопоставленный в словаре, будет скопирован в соответствующее место в иерархии модулей GraphModule. -
graph (Graph) –
graphсодержит узлы, которые должен использовать GraphModule для генерации кода -
class_name (str) –
nameобозначает имя этого GraphModule для отладки. Если оно не задано, все сообщения об ошибках будут сообщать об их происхождении изGraphModule. Может быть полезно установить это на исходное имяrootили имя, которое имеет смысл в контексте вашего преобразования.
-
root (Union[torch.nn.Module, Dict[str, Any]) –
Примечание
Обратная совместимость для этого API гарантирована.
-
add_submodule(target, m)[source] -
Добавляет указанный подмодуль в
self.Это устанавливает пустые модули там, где их ещё нет, если они являются подпутями
target.- Параметры
- Возвращает
-
- Является ли вставка подмодуля успешной. Для
-
того, чтобы этот метод возвращал True, каждый объект в цепочке, обозначенной
target, должен либо a) ещё не существовать, либо b) ссылаться наnn.Module(не параметр и не другой атрибут)
- Тип возвращаемого значения
Примечание
Обратная совместимость для этого API гарантирована.
-
property code: str -
Возвращает Python-код, сгенерированный из
Graphлежащего в основе этогоGraphModule.
-
delete_all_unused_submodules()[source] -
Удаляет все неиспользуемые подмодули из
self.Модуль считается «используемым», если истинно хотя бы одно из следующих условий: 1. У него есть дочерние модули, которые используются 2. Его forward вызывается напрямую через узел
call_module3. У него есть атрибут, не являющийся модулем, который используется узломget_attrЭтот метод можно вызвать для очистки
nn.Moduleбез ручного вызоваdelete_submoduleдля каждого неиспользуемого подмодуля.Примечание
Обратная совместимость для этого API гарантирована.
-
delete_submodule(target)[source] -
Удаляет указанный подмодуль из
self.Модуль не будет удалён, если
targetне является допустимой целью.- Параметры
-
target (str) – Полное квалифицированное строковое имя нового подмодуля (см. пример в
nn.Module.get_submoduleдля указания полного квалифицированного имени). - Возвращает
-
- Будет ли удалён подмодуль, на который ссылается заданная строка.
-
Возвращаемое значение
Falseозначает, чтоtargetне является допустимой ссылкой на подмодуль.
- Тип возвращаемого значения
Примечание
Обратная совместимость для этого API гарантирована.
-
property graph: Graph -
Возвращает
Graphлежащий в основе этогоGraphModule.
-
print_readable(print_output=True)[source] -
Возвращает Python-код, сгенерированный для текущего GraphModule и его дочерних GraphModule.
Предупреждение
Этот API является экспериментальным и НЕ обратной совместимым.
-
recompile()[source] -
Перекомпилирует этот GraphModule из его атрибута
graph. Это необходимо вызвать после редактирования содержащегосяgraph, в противном случае сгенерированный код этогоGraphModuleбудет устаревшим.Примечание
Обратная совместимость для этого API гарантирована.
- Тип возвращаемого значения
-
PythonCode
-
to_folder(folder, module_name='FxModule')[source] -
-
Dumps out module to folder with module_name so that it can be -
импортированный с
from <folder> import <module_name>Аргументы:
folder (Union[str, os.PathLike]): Папка для записи кода
-
module_name (str): Top-level name to use for the Module while -
запись кода
-
Предупреждение
Этот API является экспериментальным и НЕ обратной совместимым.
-
-
-
class torch.fx.Graph(owning_module=None, tracer_cls=None, tracer_extras=None)[source] -
Graphявляется основной структурой данных, используемой в промежуточном представлении FX. Она состоит из последовательностиNode, каждая из которых представляет вызов (или другие синтаксические конструкции). СписокNode, вместе взятый, образует допустимую функцию Python.Например, следующий код
import torch import torch.fx class MyModule(torch.nn.Module): def __init__(self): super().__init__() self.param = torch.nn.Parameter(torch.rand(3, 4)) self.linear = torch.nn.Linear(4, 5) def forward(self, x): return torch.topk(torch.sum(self.linear(x + self.linear.weight).relu(), dim=-1), 3) m = MyModule() gm = torch.fx.symbolic_trace(m)Сгенерирует следующий граф:
print(gm.graph)
graph(x): %linear_weight : [num_users=1] = self.linear.weight %add_1 : [num_users=1] = call_function[target=operator.add](args = (%x, %linear_weight), kwargs = {}) %linear_1 : [num_users=1] = call_module[target=linear](args = (%add_1,), kwargs = {}) %relu_1 : [num_users=1] = call_method[target=relu](args = (%linear_1,), kwargs = {}) %sum_1 : [num_users=1] = call_function[target=torch.sum](args = (%relu_1,), kwargs = {dim: -1}) %topk_1 : [num_users=1] = call_function[target=torch.topk](args = (%sum_1, 3), kwargs = {}) return topk_1Для семантики операций, представленных в
Graph, см.Node.Примечание
Обратная совместимость для этого API гарантирована.
-
__init__(owning_module=None, tracer_cls=None, tracer_extras=None)[source] -
Создает пустой граф.
Примечание
Обратная совместимость для этого API гарантирована.
-
call_function(the_function, args=None, kwargs=None, type_expr=None)[source] -
Вставляет
call_functionNodeвGraph. Узелcall_functionпредставляет собой вызов Python-вызываемого объекта, указанногоthe_function.- Параметры
-
-
the_function (Callable[..., Any]) – Функция, которая должна быть вызвана. Может быть любым оператором PyTorch, Python-функцией или членом
builtinsилиoperatorпространств имён. - args (Optional[Tuple[Argument, ...]]) – Позиционные аргументы, которые должны быть переданы вызываемой функции.
- kwargs (Optional[Dict[str, Argument]]) – Именованные аргументы, которые должны быть переданы вызываемой функции
- type_expr (Optional[Any]) – необязательная аннотация типа, представляющая Python-тип, который будет иметь результат этого узла.
-
the_function (Callable[..., Any]) – Функция, которая должна быть вызвана. Может быть любым оператором PyTorch, Python-функцией или членом
- Возвращает
-
Новый созданный и вставленный узел
call_function. - Тип возвращаемого значения
Примечание
Те же правила вставки и выражения типа применяются для этого метода, что и для
Graph.create_node().Примечание
Обратная совместимость для этого API гарантирована.
-
call_method(method_name, args=None, kwargs=None, type_expr=None)[source] -
Вставляет
call_methodNodeвGraph. Узелcall_methodпредставляет собой вызов заданного метода для 0-го элементаargs.- Параметры
-
-
method_name (str) – Название метода, который нужно применить к аргументу self. Например, если args[0] является
Node, представляющимTensor, то для вызоваrelu()на этомTensor, передайтеreluвmethod_name. -
args (Optional[Tuple[Argument, ...]]) – Позиционные аргументы, которые должны быть переданы вызываемому методу. Обратите внимание, что это *должно* включать аргумент
self. - kwargs (Optional[Dict[str, Argument]]) – Именованные аргументы, которые должны быть переданы вызываемому методу
- type_expr (Optional[Any]) – необязательная аннотация типа, представляющая Python-тип, который будет иметь результат этого узла.
-
method_name (str) – Название метода, который нужно применить к аргументу self. Например, если args[0] является
- Возвращает
-
Новый созданный и вставленный узел
call_method. - Тип возвращаемого значения
Примечание
Те же правила вставки и выражения типа применяются для этого метода, что и для
Graph.create_node().Примечание
Обратная совместимость для этого API гарантирована.
-
call_module(module_name, args=None, kwargs=None, type_expr=None)[source] -
Вставляет
call_moduleNodeвGraph. Узелcall_moduleпредставляет вызов функции forward() дляModuleв иерархииModule.- Параметры
-
-
module_name (str) – Полное имя
Moduleв иерархииModuleдля вызова. Например, если отслеживаемыйModuleимеет подмодуль с именемfoo, который имеет подмодуль с именемbar, полное имяfoo.barдолжно быть передано какmodule_nameдля вызова этого модуля. -
args (Optional[Tuple[Argument, ...]]) – Позиционные аргументы, которые должны быть переданы вызываемому методу. Обратите внимание, что это *не* должно включать аргумент
self. - kwargs (Optional[Dict[str, Argument]]) – Именованные аргументы, которые должны быть переданы вызываемому методу
- type_expr (Optional[Any]) – необязательная аннотация типа, представляющая Python-тип, который будет иметь результат этого узла.
-
module_name (str) – Полное имя
- Возвращает
-
Новый созданный и вставленный узел
call_module. - Тип возвращаемого значения
Примечание
Те же правила вставки и выражения типа применяются для этого метода, что и для
Graph.create_node().Примечание
Обратная совместимость для этого API гарантирована.
-
create_node(op, target, args=None, kwargs=None, name=None, type_expr=None)[source] -
Создает
Nodeи добавляет его вGraphв текущей точке вставки. Обратите внимание, что текущую точку вставки можно установить с помощьюGraph.inserting_before()иGraph.inserting_after().- Параметры
-
-
op (str) – код операции для этого узла. Один из 'call_function', 'call_method', 'get_attr', 'call_module', 'placeholder' или 'output'. Семантика этих кодов операций описана в
Graphстроке документации. - args (Optional[Tuple[Argument, ...]]) – кортеж аргументов этого узла.
- kwargs (Optional[Dict[str, Argument]]) – именованные аргументы этого узла
-
name (Optional[str]) – необязательное строковое имя для
Node. Это повлияет на имя значения, присвоенного в сгенерированном коде Python. - type_expr (Optional[Any]) – необязательная аннотация типа, представляющая Python-тип, который будет иметь результат этого узла.
-
op (str) – код операции для этого узла. Один из 'call_function', 'call_method', 'get_attr', 'call_module', 'placeholder' или 'output'. Семантика этих кодов операций описана в
- Возвращает
-
Новый созданный и вставленный узел.
- Тип возвращаемого значения
Примечание
Обратная совместимость для этого API гарантирована.
-
-
eliminate_dead_code()[source] -
Удаляет весь неиспользуемый код из графа, исходя из количества пользователей каждого узла и наличия у узлов побочных эффектов. Граф должен быть отсортирован топологически перед вызовом.
- Возвращает
-
Указывает, был ли граф изменён в результате обработки.
- Тип возвращаемого значения
Пример:
Перед удалением неиспользуемого кода
aизa = x + 1ниже не имеет пользователей и, таким образом, может быть удалён из графа без последствий.def forward(self, x): a = x + 1 return x + self.attr_1После удаления неиспользуемого кода
a = x + 1был удалён, и остальная частьforwardостаётся.def forward(self, x): return x + self.attr_1Предупреждение
Удаление неиспользуемого кода использует некоторые эвристики для предотвращения удаления узлов с побочными эффектами (см. Node.is_impure), но в целом охват очень плохой, поэтому следует предполагать, что этот метод некорректен для вызова, если вы не знаете, что ваш FX-граф состоит только из функциональных операций.
Примечание
Обратная совместимость для этого API гарантируется.
-
erase_node(to_erase)[source] -
Удаляет узел
NodeизGraph. Выбрасывает исключение, если вGraphвсё ещё есть пользователи этого узла.- Параметры
-
to_erase (Node) – Узел
Nodeдля удаления изGraph.
Примечание
Обратная совместимость для этого API гарантируется.
-
get_attr(qualified_name, type_expr=None)[source] -
Вставляет узел
get_attrв граф. Узелget_attrNodeпредставляет получение атрибута из иерархииModule.- Параметры
-
-
qualified_name (str) – полное имя атрибута, который нужно получить. Например, если отслеживаемый модуль имеет подмодуль, названный
foo, у которого есть подмодуль, названныйbar, у которого есть атрибут, названныйbaz, квалифицированное имяfoo.bar.bazдолжно быть передано какqualified_name. - type_expr (Optional[Any]) – необязательная аннотация типа, представляющая тип Python, который будет иметь выходной результат этого узла.
-
qualified_name (str) – полное имя атрибута, который нужно получить. Например, если отслеживаемый модуль имеет подмодуль, названный
- Возвращает
-
Новый созданный и вставленный узел
get_attr. - Тип возвращаемого значения
Примечание
Для этого метода применяются те же правила вставки и выражения типа, что и для
Graph.create_node.Примечание
Обратная совместимость для этого API гарантируется.
-
graph_copy(g, val_map, return_output_node=False)[source] -
Копирует все узлы из заданного графа в
self.- Параметры
- Возвращает
-
Значение в
self, которое теперь эквивалентно выходному значению вg, еслиgимел узелoutput.Noneв противном случае. - Тип возвращаемого значения
-
Optional[Union[Tuple[Any, …], List[Any], Dict[str, Any], slice, range, Node, str, int, float, bool, complex, dtype, Tensor, device, memory_format, layout, OpOverload]]
Примечание
Обратная совместимость для этого API гарантируется.
-
inserting_after(n=None)[source] -
- Устанавливает точку, в которой методы create_node и связанные с ним методы будут вставлять в граф.
-
При использовании в операторе with это временно установит точку вставки, а затем восстановит её по выходе из оператора with:
with g.inserting_after(n): ... # inserting after node n ... # insert point restored to what it was previously g.inserting_after(n) # set the insert point permanentlyАргументы:
- n (Optional[Node]): Узел, перед которым нужно вставить. Если None, то вставка произойдёт после
-
начала всего графа.
- Возвращает:
-
Управляющий ресурс, который восстановит точку вставки при выходе из
__exit__.
Примечание
Обратная совместимость для этого API гарантируется.
-
inserting_before(n=None)[source] -
- Устанавливает точку, в которой методы create_node и связанные с ним методы будут вставлять в граф.
-
При использовании в операторе with это временно установит точку вставки, а затем восстановит её по выходе из оператора with:
with g.inserting_before(n): ... # inserting before node n ... # insert point restored to what it was previously g.inserting_before(n) # set the insert point permanentlyАргументы:
- n (Optional[Node]): Узел, перед которым нужно вставить. Если None, то вставка произойдёт перед
-
началом всего графа.
- Возвращает:
-
Управляющий ресурс, который восстановит точку вставки при выходе из
__exit__.
Примечание
Обратная совместимость для этого API гарантируется.
-
lint()[source] -
Выполняет различные проверки этого графа, чтобы убедиться, что он правильно сформирован. В частности: - Проверяет, что узлы имеют правильную принадлежность (принадлежат этому графу) - Проверяет, что узлы появляются в топологическом порядке - Если у этого графа есть владелец GraphModule, проверяет, что цели существуют в этом GraphModule
Примечание
Обратная совместимость для этого API гарантируется.
-
-
node_copy(node, arg_transform=<function Graph.<lambda>>)[source] -
Копирование узла из одной графы в другую.
arg_transformнеобходимо преобразовать аргументы из графы узла в графу self. Пример:# Copying all the nodes in `g` into `new_graph` g : torch.fx.Graph = ... new_graph = torch.fx.graph() value_remap = {} for node in g.nodes: value_remap[node] = new_graph.node_copy(node, lambda n : value_remap[n])- Параметры
- Тип возвращаемого значения
Примечание
Обратная совместимость этого API гарантируется.
-
property nodes: _node_list -
Получение списка узлов, составляющих эту графу.
Обратите внимание, что это
Nodeпредставление списка — это двусвязный список. Изменения во время итерации (например, удаление узла, добавление узла) безопасны.- Возвращаемое значение
-
Двусвязный список узлов. Обратите внимание, что
reversedможет быть вызван для изменения порядка итерации.
-
on_generate_code(make_transformer)[source] -
Регистрация функции преобразования при генерации Python-кода
- Аргументы:
-
- make_transformer (Callable[[Optional[TransformCodeFunc]], TransformCodeFunc]):
-
Функция, возвращающая преобразователь кода для регистрации. Эта функция вызывается
on_generate_codeдля получения преобразователя кода.Эта функция также получает в качестве входных данных текущий зарегистрированный преобразователь кода (или None, если ничего не зарегистрировано), в случае, если не нужно его перезаписывать. Это полезно для объединения преобразователей кода.
- Возвращаемое значение:
-
менеджер контекста, который при использовании в
withоператоре автоматически восстанавливает ранее зарегистрированный преобразователь кода.
Пример:
gm: fx.GraphModule = ... # This is a code transformer we want to register. This code # transformer prepends a pdb import and trace statement at the very # beginning of the generated torch.fx code to allow for manual # debugging with the PDB library. def insert_pdb(body): return ["import pdb; pdb.set_trace()\n", *body] # Registers `insert_pdb`, and overwrites the current registered # code transformer (given by `_` to the lambda): gm.graph.on_generate_code( lambda _: insert_pdb ) # Or alternatively, registers a code transformer which first # runs `body` through existing registered transformer, then # through `insert_pdb`: gm.graph.on_generate_code( lambda current_trans: ( lambda body: insert_pdb( current_trans(body) if current_trans else body ) ) ) gm.recompile() gm(*inputs) # drops into pdbЭта функция также может использоваться в качестве менеджера контекста с преимуществом автоматического восстановления ранее зарегистрированного преобразователя кода:
# ... continue from previous example with gm.graph.on_generate_code(lambda _: insert_pdb): # do more stuff with `gm`... gm.recompile() gm(*inputs) # drops into pdb # now previous code transformer is restored (but `gm`'s code with pdb # remains - that means you can run `gm` with pdb here too, until you # run next `recompile()`).Предупреждение
Этот API является экспериментальным и НЕ обратной совместимым.
-
output(result, type_expr=None)[source] -
Вставка
outputNodeвGraph. Узелoutputпредставляет операторreturnв Python-коде.result— это значение, которое должно быть возвращено.- Параметры
-
- result (Argument) – Возвращаемое значение.
- type_expr (Optional[Any]) – необязательная аннотация типа, представляющая Python-тип, который будет у результата этого узла.
Примечание
Те же правила вставки и выражения типа применяются для этого метода, что и для
Graph.create_node.Примечание
Обратная совместимость этого API гарантирована.
-
placeholder(name, type_expr=None, default_value)[source] -
Вставка узла
placeholderв графу. Узелplaceholderпредставляет вход функции.- Параметры
-
-
name (str) – Имя входного значения. Соответствует имени позиционного аргумента функции, которую представляет этот
Graph. - type_expr (Optional[Any]) – необязательная аннотация типа, представляющая Python-тип, который будет у результата этого узла. Это необходимо в некоторых случаях для правильной генерации кода (например, при последующем использовании функции в TorchScript).
-
default_value (Any) – Значение по умолчанию для этого аргумента функции. ПРИМЕЧАНИЕ: чтобы
Noneможно было использовать в качестве значения по умолчанию,inspect.Signature.emptyдолжно быть передано в качестве этого аргумента, чтобы указать, что параметр _не_ имеет значения по умолчанию.
-
name (str) – Имя входного значения. Соответствует имени позиционного аргумента функции, которую представляет этот
- Тип возвращаемого значения
Примечание
Те же правила вставки и выражения типа применяются для этого метода, что и для
Graph.create_node.Примечание
Обратная совместимость этого API гарантирована.
-
print_tabular()[source] -
Вывод промежуточного представления графика в табличной форме. Обратите внимание, что для этого API требуется установка модуля
tabulate.Примечание
Обратная совместимость этого API гарантирована.
-
process_inputs(*args)[source] -
Обработка аргументов для их передачи в граф FX.
Предупреждение
Этот API является экспериментальным и НЕ обратной совместимым.
-
process_outputs(out)[source] -
Предупреждение
Этот API является экспериментальным и НЕ обратной совместимым.
-
python_code(root_module, *, verbose=False)[source] -
Преобразование этой
Graphв действительный Python-код.- Параметры
-
root_module (str) – Имя корневого модуля, в котором будут искать целевые имена.
- Возвращаемое значение
-
src: исходный Python-код, представляющий объект globals: словарь глобальных имен в
src-> соответствующие им объекты. - Тип возвращаемого значения
-
Объект PythonCode, состоящий из двух полей
Примечание
Обратная совместимость этого API гарантирована.
-
set_codegen(codegen)[source] -
Предупреждение
Этот API является экспериментальным и НЕ обратной совместимым.
-
-
class torch.fx.Node(graph, name, op, target, args, kwargs, return_type=None)[source] -
Node— это структура данных, представляющая отдельные операции вGraph. В основном, узлы представляют вызовы различных сущностей, таких как операторы, методы и модули (некоторые исключения включают узлы, определяющие входные и выходные данные функции). У каждогоNodeуказана функция, определённая свойствомop. СемантикаNodeдля каждого значенияopследующие:-
placeholderпредставляет вход функции. Атрибутnameопределяет имя, которое получит это значение.targetаналогичным образом — имя аргумента.argsсодержит либо: 1) ничего, либо 2) один аргумент, обозначающий параметр по умолчанию для входных данных функции.kwargs— это игнорируемое значение. Заполнители соответствуют параметрам функции (например,x) в выводе графа. -
get_attrизвлекает параметр из иерархии модулей.nameаналогичным образом — имя, присваиваемое результату извлечения.target— полностью квалифицированное имя позиции параметра в иерархии модулей.argsиkwargs— это игнорируемые значения. -
call_functionприменяет свободную функцию к некоторым значениям.nameаналогичным образом — имя присваиваемого значения.target— это функция, которая должна быть применена.argsиkwargsпредставляют аргументы функции, следуя соглашениям Python для вызовов функций. -
call_moduleприменяет методforward()модуля в иерархии модулей к заданным аргументам.name— как и раньше.target— это полностью квалифицированное имя модуля в иерархии модулей для вызова.argsиkwargsпредставляют аргументы для вызова модуля, *исключая аргумент self*. -
call_methodвызывает метод на значении.name— как и раньше.target— строковое имя метода, который нужно применить к аргументуself.argsиkwargsпредставляют аргументы для вызова модуля, *включая аргумент self*. -
outputсодержит результат прослеживаемой функции в атрибутеargs[0]. Это соответствует инструкции «return» в выводе графа.
Примечание
Обратная совместимость для этого API гарантирована.
-
property all_input_nodes: List[Node] -
Возвращает все узлы, являющиеся входными для этого узла. Это эквивалентно перебору
argsиkwargsи сбору только тех значений, которые являются узлами.- Возвращает
-
Список
Nodes, которые появляются вargsиkwargsэтогоNode, в указанном порядке.
-
append(x)[source] -
Вставляет
xпосле этого узла в список узлов в графе. Эквивалентноself.next.prepend(x)- Параметры
-
x (Узел) – Узел, который следует поместить после этого узла. Должен быть членом того же графа.
Примечание
Обратная совместимость для этого API гарантирована.
-
property args: Tuple[Optional[Union[Tuple[Any, ...], List[Any], Dict[str, Any], slice, range, Node, str, int, float, bool, complex, dtype, Tensor, device, memory_format, layout, OpOverload]], ...] -
Кортеж аргументов для этого
Node. Интерпретация аргументов зависит от кода операции узла. Дополнительную информацию см. в документацииNode.Присваивание этому свойству разрешено. Все учёт использования и пользователей обновляется автоматически при присваивании.
-
format_node(placeholder_names=None, maybe_return_typename=None)[source] -
Возвращает строковое описание узла
self. Это метод может использоваться без аргументов в качестве инструмента отладки.Этот метод также используется внутри метода
__str__Graph. Вместе строки вplaceholder_namesиmaybe_return_typenameсоставляют подпись автоматически генерируемой функцииforwardв окружающем GraphModule этого графа.placeholder_namesиmaybe_return_typenameне должны использоваться в иных случаях.- Параметры
-
-
placeholder_names (Optional[Список[строка]]) – Список, который будет хранить отформатированные строки, представляющие заполнители в сгенерированной функции
forward. Только внутреннее использование. -
maybe_return_typename (Optional[Список[строка]]) – Одноэлементный список, который будет хранить отформатированную строку, представляющую результат сгенерированной функции
forward. Только внутреннее использование.
-
placeholder_names (Optional[Список[строка]]) – Список, который будет хранить отформатированные строки, представляющие заполнители в сгенерированной функции
- Возвращает
-
-
If 1) we’re using format_node as an internal helper -
1) в методе
__str__Graph, и 2) еслиselfявляется замещающим узлом, возвращаетNone. В противном случае возвращает описательную строку представления текущего узла.
-
- Тип возвращаемого значения
Примечание
Обратная совместимость для этого API гарантирована.
-
is_impure()[source] -
Возвращает, является ли эта операция нечистой, т.е. если её операция является заполнителем или выходом, или если вызов функции или вызов модуля нечистый.
- Возвращает
-
Является ли операция нечистой.
- Тип возвращаемого значения
Предупреждение
Этот API экспериментальный и не гарантирует обратную совместимость.
-
property kwargs: Dict[str, Optional[Union[Tuple[Any, ...], List[Any], Dict[str, Any], slice, range, Node, str, int, float, bool, complex, dtype, Tensor, device, memory_format, layout, OpOverload]]] -
Словарь ключевых аргументов для этого
Node. Интерпретация аргументов зависит от кода операции узла. Дополнительную информацию см. в документацииNode.Присваивание этому свойству разрешено. Все учёт использования и пользователей обновляется автоматически при присваивании.
-
property next: Node -
Возвращает следующий
Nodeв связанном списке узлов.- Возвращает
-
Следующий
Nodeв связанном списке узлов.
-
-
normalized_arguments(root, arg_types=None, kwarg_types=None, normalize_to_only_use_kwargs=False)[source] -
Возвращает нормализованные аргументы для Python-целей. Это означает, что
args/kwargsбудут сопоставлены с сигнатурой модуля/функции и вернут исключительно аргументы в порядке следования, еслиnormalize_to_only_use_kwargsимеет значение true. Также заполняются значения по умолчанию. Не поддерживает позиционные-только параметры или параметры varargs.Поддерживает вызовы модулей.
Возможно, потребуется
arg_typesиkwarg_typesдля разбора перегрузок.- Параметры
-
- root (torch.nn.Module) – Модуль, на котором необходимо разрешить модульные цели.
- arg_types (Optional[Tuple[Any]]) – Кортеж типов аргументов для аргументов
- kwarg_types (Optional[Dict[str, Any]]) – Словарь типов аргументов для ключевых аргументов
- normalize_to_only_use_kwargs (bool) – Принудительное использование только ключевых аргументов.
- Возвращает
-
Возвращает кортеж ArgsKwargsPair, или
Noneв случае неудачи. - Тип возвращаемого значения
-
Optional[ArgsKwargsPair]
Предупреждение
Этот API находится в стадии разработки и НЕ совместим с предыдущими версиями.
-
prepend(x)[source] -
Вставить x перед этим узлом в список узлов в графе. Пример:
Before: p -> self bx -> x -> ax After: p -> x -> self bx -> ax- Параметры
-
x (Узел) – Узел, который нужно поместить перед этим узлом. Должен быть членом той же схемы.
Примечание
Обратная совместимость этого API гарантирована.
-
property prev: Node -
Возвращает предыдущий
Nodeв связанном списке узлов.- Возвращает
-
Предыдущий
Nodeв связанном списке узлов.
-
replace_all_uses_with(replace_with, delete_user_cb=<function Node.<lambda>>, *, propagate_meta=False)[source] -
Заменить все использования
selfв схеме на узелreplace_with.- Параметры
-
-
replace_with (Узел) – Узел, на который нужно заменить все использования
self. - delete_user_cb (Callable) – Обратный вызов, который вызывается для определения, следует ли удалять данного пользователя узла self.
- propagate_meta (bool) – Нужно ли копировать все свойства в поле .meta исходного узла на узел-замену. Для безопасности это допустимо только в том случае, если у узла-замены нет существующего поля .meta.
-
replace_with (Узел) – Узел, на который нужно заменить все использования
- Возвращает
-
Список узлов, на которых было произведено это изменение.
- Тип возвращаемого значения
Примечание
Обратная совместимость этого API гарантирована.
-
replace_input_with(old_input, new_input)[source] -
Перебрать входные узлы
self, и заменить все экземплярыold_inputнаnew_input.- Параметры
Примечание
Обратная совместимость этого API гарантирована.
-
property stack_trace: Optional[str] -
Возвращает трассировку стека Python, записанную во время трассировки, если таковая имеется. При трассировке с помощью fx.Tracer это свойство обычно заполняется
Tracer.create_proxy. Чтобы записывать трассировки стека во время трассировки для отладки, установитеrecord_stack_traces = Trueв экземпляреTracer. При трассировке с помощью dynamo это свойство будет заполняться по умолчаниюOutputGraph.create_proxy.stack_trace будет содержать внутренний кадр в конце строки.
-
update_arg(idx, arg)[source] -
Обновить существующий позиционный аргумент, чтобы он содержал новое значение
arg. После вызоваself.args[idx] == arg.- Параметры
-
-
idx (int) – Индекс в
self.argsдля обновления элемента -
arg (Argument) – Новое значение аргумента, которое нужно записать в
args
-
idx (int) – Индекс в
Примечание
Обратная совместимость этого API гарантирована.
-
update_kwarg(key, arg)[source] -
Обновить существующий ключевой аргумент, чтобы он содержал новое значение
arg. После вызоваself.kwargs[key] == arg.- Параметры
-
-
key (str) – Ключ в
self.kwargsдля обновления элемента -
arg (Argument) – Новое значение аргумента, которое нужно записать в
kwargs
-
key (str) – Ключ в
Примечание
Обратная совместимость этого API гарантирована.
-
-
class torch.fx.Tracer(autowrap_modules=(math,), autowrap_functions=())[source] -
Tracer— это класс, реализующий функциональность символического трассирования вtorch.fx.symbolic_trace. Вызовsymbolic_trace(m)эквивалентен вызовуTracer().trace(m).Класс Tracer можно наследоваться, чтобы переопределить различные аспекты процесса трассирования. Различные переопределяемые аспекты описаны в документации методов данного класса.
Примечание
Обратная совместимость данного API гарантируется.
-
call_module(m, forward, args, kwargs)[source] -
Метод, определяющий поведение данного
Tracerпри встрече вызова экземпляраnn.Module.По умолчанию, поведение заключается в проверке, является ли вызываемый модуль листовым модулем с помощью
is_leaf_module. Если да, то генерируется узелcall_module, ссылающийся наmвGraph. В противном случае, вызываетсяModuleстандартным образом, прослеживая операции в его методеforward.Этот метод можно переопределить, например, для создания вложенных отслеживаемых GraphModules или для любого другого поведения, необходимого при трассировании через границы
Module.- Параметры
-
- m (Модуль) — Модуль, для которого генерируется вызов.
-
forward (Callable) — Метод forward() вызываемого
Module. - args (Tuple) — Аргументы вызова модуля.
- kwargs (Dict) — Имена аргументов вызова модуля.
- Возвращает
-
Значение, возвращаемое вызовом модуля. В случае генерации узла
call_module, это значение типаProxy. В противном случае, это значение, возвращённое вызовомModule. - Тип возвращаемого значения
Примечание
Обратная совместимость данного API гарантируется.
-
create_arg(a)[source] -
Метод, определяющий поведение трассирования при подготовке значений для использования в качестве аргументов узлов в
Graph.По умолчанию, поведение включает:
- Итерацию по коллекционным типам (например, кортеж, список, словарь) и рекурсивный вызов
create_argsдля элементов. - Для объекта Proxy возвращает ссылку на подлежащий IR
Node. -
Для объекта тензора, не являющегося Proxy, генерирует IR для различных случаев:
- Для параметра генерирует узел
get_attr, ссылающийся на этот параметр. - Для тензора, не являющегося параметром, сохраняет тензор в специальном атрибуте, ссылающемся на этот атрибут.
- Для параметра генерирует узел
Этот метод можно переопределить для поддержки других типов.
- Параметры
-
a (Any) — Значение, которое должно быть сгенерировано как
ArgumentвGraph. - Возвращает
-
Значение
a, преобразованное в соответствующийArgument. - Тип возвращаемого значения
-
Optional[Union[Tuple[Any, …], List[Any], Dict[str, Any], slice, range, Node, str, int, float, bool, complex, dtype, Tensor, device, memory_format, layout, OpOverload]]
Примечание
Обратная совместимость данного API гарантируется.
- Итерацию по коллекционным типам (например, кортеж, список, словарь) и рекурсивный вызов
-
create_args_for_root(root_fn, is_module, concrete_args=None)[source] -
Создаёт узлы
placeholderсоответствующие сигнатуре модуляroot. Данный метод интроспектирует сигнатуру root и генерирует соответствующие узлы, также поддерживая*argsи**kwargs.Предупреждение
Данный API экспериментальный и НЕ обратной совместимости.
-
create_node(kind, target, args, kwargs, name=None, type_expr=None) -
Вставляет узел графа, используя target, args, kwargs и имя.
Этот метод можно переопределить для выполнения дополнительной проверки, валидации или модификации значений, используемых при создании узла. Например, можно запретить запись операций на месте.
Примечание
Обратная совместимость данного API гарантируется.
- Тип возвращаемого значения
-
create_proxy(kind, target, args, kwargs, name=None, type_expr=None, proxy_factory_fn=None) -
Создаёт узел из заданных аргументов, затем возвращает узел, обернутый в объект Proxy.
Если kind = ‘placeholder’, то мы создаём узел, представляющий параметр функции. Если нам нужно закодировать параметр по умолчанию, мы используем кортеж
args.argsв противном случае пуст для узловplaceholder.Примечание
Обратная совместимость данного API гарантируется.
-
getattr(attr, attr_val, parameter_proxy_cache)[source] -
Метод, определяющий поведение данного
Tracerпри вызове getattr для экземпляраnn.Module.По умолчанию, поведение заключается в возвращении значения proxy для атрибута. Оно также сохраняет значение proxy в
parameter_proxy_cache, поэтому последующие вызовы будут повторно использовать proxy, а не создавать новый.Этот метод можно переопределить, например, для того, чтобы не возвращать proxy при запросе параметров.
- Параметры
- Возвращает
-
Значение, возвращаемое вызовом getattr.
Предупреждение
Данный API экспериментальный и НЕ обратной совместимости.
-
-
is_leaf_module(m, module_qualified_name)[source] -
Метод для определения, является ли данный
nn.Module«листовым» модулем.Листовые модули являются атомными единицами, которые появляются в IR, на которые ссылаются вызовы
call_module. По умолчанию, модули в пространстве имён стандартной библиотеки PyTorch (torch.nn) являются листовыми модулями. Все остальные модули отслеживаются, и их составляющие операции записываются, если не указано иное с помощью этого параметра.- Параметры
- Тип возвращаемого значения
Примечание
Совместимость с предыдущими версиями для этого API гарантирована.
-
iter(obj) -
- Вызывается при итерировании по объекту-прокси, например,
-
при использовании в управляющих структурах. Обычно мы не знаем, что делать, потому что не знаем значение прокси, но пользовательский трейсер может добавить больше информации к узлу графа с помощью create_node и выбрать возвращение итератора.
Примечание
Совместимость с предыдущими версиями для этого API гарантирована.
- Тип возвращаемого значения
-
keys(obj) -
- Вызывается при вызове метода keys() у объекта-прокси.
-
Это происходит при вызове ** у прокси. Это должно вернуть итератор, если ** должен работать в вашем пользовательском трейсере.
Примечание
Совместимость с предыдущими версиями для этого API гарантирована.
- Тип возвращаемого значения
-
path_of_module(mod)[source] -
Вспомогательный метод для поиска квалифицированного имени
modв иерархии модулейroot. Например, если уrootесть подмодуль с именемfoo, который имеет подмодуль с именемbar, передачаbarв эту функцию вернёт строку “foo.bar”.- Параметры
-
mod (строка) – Модуль для которого необходимо получить квалифицированное имя.
- Тип возвращаемого значения
Примечание
Совместимость с предыдущими версиями для этого API гарантирована.
-
proxy(node) -
Примечание
Совместимость с предыдущими версиями для этого API гарантирована.
- Тип возвращаемого значения
-
to_bool(obj) -
- Вызывается при преобразовании объекта-прокси в булевое значение, например,
-
при использовании в управляющих структурах. Обычно мы не знаем, что делать, потому что не знаем значение прокси, но пользовательский трейсер может добавить больше информации к узлу графа с помощью create_node и выбрать возвращение значения.
Примечание
Совместимость с предыдущими версиями для этого API гарантирована.
- Тип возвращаемого значения
-
trace(root, concrete_args=None)[source] -
Отследить
rootи вернуть соответствующее представление FXGraph.rootможет быть экземпляромnn.Moduleили Python-функцией.Обратите внимание, что после этого вызова
self.rootможет отличаться отroot, переданного сюда. Например, когда свободной функции передаетсяtrace(), мы создадим экземплярnn.Moduleв качестве корня и добавим встроенные константы.- Параметры
-
-
root (Union[Модуль, Callable]) – Либо
Module, либо функция, которая должна быть отслежена. Совместимость с предыдущими версиями для этого параметра гарантируется. - concrete_args (Optional[Dict[строка, любой]]) – Конкретные аргументы, которые не должны обрабатываться как прокси. Этот параметр экспериментальный, и его обратная совместимость НЕ гарантируется.
-
root (Union[Модуль, Callable]) – Либо
- Возвращает
-
FX-граф, представляющий семантику переданного
root. - Тип возвращаемого значения
Примечание
Совместимость с предыдущими версиями для этого API гарантирована.
-
-
class torch.fx.Proxy(node, tracer=None)[source] -
Объекты
ProxyявляютсяNodeобёртками, которые проходят через программу во время символического отслеживания и записывают все операции (torchвызовы функций, вызовы методов, операторы), с которыми они взаимодействуют, в растущий FX-граф.Если вы выполняете преобразования графа, вы можете обернуть свой собственный метод
Proxyвокруг исходногоNode, чтобы использовать перегруженные операторы для добавления дополнительных элементов вGraph.Объекты
Proxyне могут быть итерированы. Другими словами, символический трейсер выдаст ошибку, еслиProxyиспользуется в цикле или в качестве аргумента функции*args/**kwargs.Есть два основных способа обойти это: 1. Вынести неотслеживаемую логику в функцию верхнего уровня и применить
fx.wrapк ней. 2. Если управляющая структура статична (т. е. количество циклов основано на некотором гиперпараметре), код можно оставить на своем месте и переработать в нечто вроде:for i in range(self.some_hyperparameter): indexed_item = proxied_value[i]Для более подробного описания внутренних механизмов прокси, обратитесь к разделу «Прокси» в
torch/fx/OVERVIEW.mdПримечание
Совместимость с предыдущими версиями для этого API гарантирована.
-
class torch.fx.Interpreter(module, garbage_collect_values=True)[source] -
Интерпретатор выполняет узел FX-графа по узлу. Этот шаблон может быть полезен для многих задач, включая написание преобразований кода и проход анализа.
Методы в классе Interpreter могут быть переопределены для настройки поведения выполнения. Карта переопределяемых методов с точки зрения иерархии вызовов:
run() +-- run_node +-- placeholder() +-- get_attr() +-- call_function() +-- call_method() +-- call_module() +-- output()Пример
Предположим, что мы хотим поменять все экземпляры
torch.negнаtorch.sigmoidи наоборот (включая их эквиваленты методаTensor). Мы можем создать подкласс Interpreter следующим образом:class NegSigmSwapInterpreter(Interpreter): def call_function(self, target : Target, args : Tuple, kwargs : Dict) -> Any: if target == torch.sigmoid: return torch.neg(*args, **kwargs) return super().call_function(n) def call_method(self, target : Target, args : Tuple, kwargs : Dict) -> Any: if target == 'neg': call_self, *args_tail = args return call_self.sigmoid(*args_tail, **kwargs) return super().call_method(n) def fn(x): return torch.sigmoid(x).neg() gm = torch.fx.symbolic_trace(fn) input = torch.randn(3, 4) result = NegSigmSwapInterpreter(gm).run(input) torch.testing.assert_close(result, torch.neg(input).sigmoid())- Параметры
-
- module (GraphModule) – Модуль, подлежащий выполнению
-
garbage_collect_values (bool) – Удалять ли значения после их последнего использования в рамках выполнения модуля. Это гарантирует оптимальное использование памяти во время выполнения. Это можно отключить, чтобы, например, просмотреть все промежуточные значения в выполнении, посмотрев на атрибут
Interpreter.env.
Примечание
Обратная совместимость для этого API гарантируется.
-
boxed_run(args_list)[source] -
Выполнить
moduleпосредством интерпретации и вернуть результат. Это использует соглашение вызова «boxed», где вы передаёте список аргументов, который будет очищен интерпретатором. Это гарантирует, что входные тензоры будут немедленно освобождены.Примечание
Обратная совместимость для этого API гарантируется.
-
call_function(target, args, kwargs)[source] -
Выполнить узел
call_functionи вернуть результат.- Параметры
-
- target (Target) – Цель вызова для данного узла. См. Node для получения подробной информации о семантике
- args (Tuple) – Кортеж позиционных аргументов для этого вызова
- kwargs (Dict) – Словарь ключевых аргументов для этого вызова
- Тип возвращаемого значения
- Возвращаемое значение
-
Any: значение, возвращённое вызовом функции
Примечание
Обратная совместимость для этого API гарантируется.
-
call_method(target, args, kwargs)[source] -
Выполнить узел
call_methodи вернуть результат.- Параметры
-
- target (Target) – Цель вызова для данного узла. См. Node для получения подробной информации о семантике
- args (Tuple) – Кортеж позиционных аргументов для этого вызова
- kwargs (Dict) – Словарь ключевых аргументов для этого вызова
- Тип возвращаемого значения
- Возвращаемое значение
-
Any: значение, возвращённое вызовом метода
Примечание
Обратная совместимость для этого API гарантируется.
-
call_module(target, args, kwargs)[source] -
Выполнить узел
call_moduleи вернуть результат.- Параметры
-
- target (Target) – Цель вызова для данного узла. См. Node для получения подробной информации о семантике
- args (Tuple) – Кортеж позиционных аргументов для этого вызова
- kwargs (Dict) – Словарь ключевых аргументов для этого вызова
- Тип возвращаемого значения
- Возвращаемое значение
-
Any: значение, возвращённое вызовом модуля
Примечание
Обратная совместимость для этого API гарантируется.
-
fetch_args_kwargs_from_env(n)[source] -
Извлечь конкретные значения
argsиkwargsузлаnиз текущей среды выполнения.- Параметры
-
n (Node) – Узел, для которого необходимо извлечь
argsиkwargs. - Возвращаемое значение
-
argsиkwargsс конкретными значениями дляn. - Тип возвращаемого значения
-
Tuple[Tuple, Dict]
Примечание
Обратная совместимость для этого API гарантируется.
-
fetch_attr(target)[source] -
Извлечь атрибут из иерархии
Moduleself.module.- Параметры
-
target (str) – Полное квалифицированное имя атрибута для извлечения
- Возвращаемое значение
-
Значение атрибута.
- Тип возвращаемого значения
-
Any
Примечание
Обратная совместимость для этого API гарантируется.
-
get_attr(target, args, kwargs)[source] -
Выполнить узел
get_attr. Извлечёт значение атрибута из иерархииModuleself.module.- Параметры
-
- target (Target) – Цель вызова для данного узла. См. Node для получения подробной информации о семантике
- args (Tuple) – Кортеж позиционных аргументов для этого вызова
- kwargs (Dict) – Словарь ключевых аргументов для этого вызова
- Возвращаемое значение
-
Значение извлечённого атрибута.
- Тип возвращаемого значения
-
Any
Примечание
Обратная совместимость для этого API гарантируется.
-
map_nodes_to_values(args, n)[source] -
Рекурсивно проходит по
argsи ищет конкретное значение для каждогоNodeв текущей среде выполнения.- Параметры
-
- args (Argument) – Структура данных для поиска конкретных значений
-
n (Узел) – Узел, к которому
argsотносится. Используется только для отладки ошибок.
- Тип возвращаемого значения
-
Optional[Union[Кортеж[Любой, …], Список[Любой], Словарь[строка, Любой], срез, диапазон, Узел, строка, целое число, вещественное число, логическое значение, комплексное число, тип данных, тензор, устройство, форма памяти, структура, OpOverload]]
Примечание
Обратная совместимость для этого API гарантируется.
-
output(target, args, kwargs)[source] -
Выполняет узел
output. Это просто извлекает значение, на которое ссылается узелoutput, и возвращает его.- Параметры
-
- target (Target) – Цель вызова для этого узла. См. Узел для получения подробной информации о семантике
- args (Tuple) – Кортеж позиционных аргументов для этого вызова
- kwargs (Dict) – Словарь ключевых аргументов для этого вызова
- Возвращает
-
Возвращаемое значение, на которое указывает выходной узел
- Тип возвращаемого значения
-
Любой
Примечание
Обратная совместимость для этого API гарантируется.
-
placeholder(target, args, kwargs)[source] -
Выполняет узел
placeholder. Обратите внимание, что это состояние:Interpreterподдерживает внутренний итератор по аргументам, переданным вrun, и этот метод возвращает next() для этого итератора.- Параметры
-
- target (Target) – Цель вызова для этого узла. См. Узел для получения подробной информации о семантике
- args (Tuple) – Кортеж позиционных аргументов для этого вызова
- kwargs (Dict) – Словарь ключевых аргументов для этого вызова
- Возвращает
-
Значение аргумента, которое было получено.
- Тип возвращаемого значения
-
Любой
Примечание
Обратная совместимость для этого API гарантируется.
-
run(*args, initial_env=None, enable_io_processing=True)[source] -
Выполняет
moduleчерез интерпретацию и возвращает результат.- Параметры
-
- *args – Аргументы модуля для выполнения в позиционном порядке
-
initial_env (Optional[Dict[Узел, Любой]]) – Необязательная начальная среда выполнения. Это словарь, сопоставляющий
Nodeлюбому значению. Это можно использовать, например, для предварительной заполнения результатов для определенныхNodesдля частичной оценки в интерпретаторе. - enable_io_processing (bool) – Если True, мы сначала обрабатываем входные и выходные данные функциями process_inputs и process_outputs графа, прежде чем использовать их.
- Возвращает
-
Значение, возвращенное при выполнении модуля
- Тип возвращаемого значения
-
Любой
Примечание
Обратная совместимость для этого API гарантируется.
-
run_node(n)[source] -
Выполняет конкретный узел
nи возвращает результат. Вызывает placeholder, get_attr, call_function, call_method, call_module или output в зависимости отnode.op- Параметры
-
n (Узел) – Узел для выполнения
- Возвращает
-
Результат выполнения
n - Тип возвращаемого значения
-
Любой
Примечание
Обратная совместимость для этого API гарантируется.
-
-
class torch.fx.Transformer(module)[source] -
Transformer— это специальный тип интерпретатора, который создаёт новыйModule. Он предоставляет методtransform(), который возвращает преобразованныйModule.Transformerне требует аргументов для выполнения, в отличие отInterpreter.Transformerработает полностью символически.Пример
Предположим, мы хотим поменять все экземпляры
torch.negнаtorch.sigmoidи наоборот (включая эквиваленты методовTensor). Мы можем сделать это, наследуя классTransformer:class NegSigmSwapXformer(Transformer): def call_function(self, target : 'Target', args : Tuple[Argument, ...], kwargs : Dict[str, Any]) -> Any: if target == torch.sigmoid: return torch.neg(*args, **kwargs) return super().call_function(n) def call_method(self, target : 'Target', args : Tuple[Argument, ...], kwargs : Dict[str, Any]) -> Any: if target == 'neg': call_self, *args_tail = args return call_self.sigmoid(*args_tail, **kwargs) return super().call_method(n) def fn(x): return torch.sigmoid(x).neg() gm = torch.fx.symbolic_trace(fn) transformed : torch.nn.Module = NegSigmSwapXformer(gm).transform() input = torch.randn(3, 4) torch.testing.assert_close(transformed(input), torch.neg(input).sigmoid())- Параметры
-
module (GraphModule) — Преобразуемый
Module.
Примечание
Обратная совместимость для этого API гарантируется.
-
call_function(target, args, kwargs)[source] -
Примечание
Обратная совместимость для этого API гарантируется.
- Тип возвращаемого значения
-
call_module(target, args, kwargs)[source] -
Примечание
Обратная совместимость для этого API гарантируется.
- Тип возвращаемого значения
-
get_attr(target, args, kwargs)[source] -
Выполнить узел
get_attr. ВTransformer, это переопределяется для вставки нового узлаget_attrв граф вывода.- Параметры
-
- target (Target) — Цель вызова для этого узла. Смотрите Node для получения подробностей о семантике
- args (Tuple) — Кортеж позиционных аргументов для этого вызова
- kwargs (Dict) — Словарь ключевых аргументов для этого вызова
- Тип возвращаемого значения
Примечание
Обратная совместимость для этого API гарантируется.
-
placeholder(target, args, kwargs)[source] -
Выполнить узел
placeholder. ВTransformer, это переопределяется для вставки новогоplaceholderв граф вывода.- Параметры
-
- target (Target) — Цель вызова для этого узла. Смотрите Node для получения подробностей о семантике
- args (Tuple) — Кортеж позиционных аргументов для этого вызова
- kwargs (Dict) — Словарь ключевых аргументов для этого вызова
- Тип возвращаемого значения
Примечание
Обратная совместимость для этого API гарантируется.
-
transform()[source] -
Преобразовать
self.moduleи вернуть преобразованныйGraphModule.Примечание
Обратная совместимость для этого API гарантируется.
- Тип возвращаемого значения
-
torch.fx.replace_pattern(gm, pattern, replacement)[source] -
Сопоставляет все возможные непересекающиеся наборы операторов и их зависимостей данных (
pattern) в графе GraphModule (gm), затем заменяет каждый из этих сопоставленных подграфов другим подграфом (replacement).- Параметры
-
- gm (GraphModule) — GraphModule, который оборачивает Graph для обработки
-
pattern (Union[Callable, GraphModule]) — Подграф для поиска в
gmдля замены -
replacement (Union[Callable, GraphModule]) — Подграф, которым нужно заменить
pattern
- Возвращаемое значение
-
Список объектов
Match, представляющих места в исходном графе, где было найдено совпадение сpattern. Список пуст, если совпадений нет.Matchопределяется как:class Match(NamedTuple): # Node from which the match was found anchor: Node # Maps nodes in the pattern subgraph to nodes in the larger graph nodes_map: Dict[Node, Node] - Тип возвращаемого значения
-
List[Match]
Примеры:
import torch from torch.fx import symbolic_trace, subgraph_rewriter class M(torch.nn.Module): def __init__(self): super().__init__() def forward(self, x, w1, w2): m1 = torch.cat([w1, w2]).sum() m2 = torch.cat([w1, w2]).sum() return x + torch.max(m1) + torch.max(m2) def pattern(w1, w2): return torch.cat([w1, w2]).sum() def replacement(w1, w2): return torch.stack([w1, w2]) traced_module = symbolic_trace(M()) subgraph_rewriter.replace_pattern(traced_module, pattern, replacement)Этот код сначала найдет совпадение с
patternв методеforwardклассаtraced_module. Поиск совпадений производится по отношениям использования и определения, а не по именам узлов. Например, если у вас естьp = torch.cat([a, b])вpattern, вы можете найти совпадение сm = torch.cat([a, b])в исходной функцииforward, несмотря на то, что имена переменных разные (pvsm).Выражение
returnвpatternищется только по своему значению; оно может или не может совпасть с выражениемreturnв более крупном графе. Другими словами, шаблон не должен распространяться до конца более крупного графа.Когда шаблон найден, он удаляется из более крупной функции и заменяется
replacement. Если в более крупной функции есть несколько совпадений сpattern, каждое непересекающееся совпадение будет заменено. В случае перекрытия совпадений, будет заменено первое найденное совпадение в наборе перекрывающихся совпадений. ("Первое" здесь определяется как первое в топологическом порядке отношений использования и определения узлов. В большинстве случаев, первый узел — это параметр, который появляется непосредственно послеself, а последний узел — то, что функция возвращает.)Важно отметить, что параметры
patternCallable должны использоваться в самом Callable, а параметрыreplacementCallable должны совпадать с шаблоном. Первый принцип объясняет, почему в приведенном коде функцияforwardимеет параметрыx, w1, w2, но функцияpatternтолькоw1, w2.patternне используетx, поэтому не должна указыватьxкак параметр. В качестве примера второго принципа рассмотрим заменуdef pattern(x, y): return torch.neg(x) + torch.relu(y)на
def replacement(x, y): return torch.relu(x)В этом случае
replacementнуждается в таком же количестве параметров, как иpattern(какx, так иy), даже если параметрyне используется вreplacement.После вызова
subgraph_rewriter.replace_patternсгенерированный Python-код выглядит так:def forward(self, x, w1, w2): stack_1 = torch.stack([w1, w2]) sum_1 = stack_1.sum() stack_2 = torch.stack([w1, w2]) sum_2 = stack_2.sum() max_1 = torch.max(sum_1) add_1 = x + max_1 max_2 = torch.max(sum_2) add_2 = add_1 + max_2 return add_2Примечание
Обратная совместимость для этого API гарантируется.
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/fx.html