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 : [#users=1] = placeholder[target=x]
%param : [#users=1] = get_attr[target=param]
%add : [#users=1] = call_function[target=operator.add](args = (%x, %param), kwargs = {})
%linear : [#users=1] = call_module[target=linear](args = (%add,), kwargs = {})
%clamp : [#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-кода. Он подаёт поддельные значения, называемые Proxy, через код. Операции с этими Proxy записываются. Дополнительную информацию о символическом прослеживании можно найти в документации symbolic_trace() и Tracer.
Промежуточное представление — это контейнер для операций, которые были записаны во время символического прослеживания. Оно состоит из списка узлов, представляющих входные данные функции, места вызова (функций, методов или torch.nn.Module-экземпляров) и возвращаемые значения. Дополнительную информацию о промежуточном представлении можно найти в документации Graph. Именно в промежуточном представлении применяются преобразования.
Генерация Python-кода делает FX инструментом для преобразования Python-кода в Python-код (или модуля в модуль). Для каждого графа промежуточного представления мы можем создать допустимый Python-код, соответствующий семантике графа. Эта функциональность обернута в GraphModule, который является экземпляром torch.nn.Module, содержащим Graph и метод forward, сгенерированный из графа.
Вместе эти компоненты (символическое прослеживание -> промежуточное представление -> преобразования -> генерация 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 | имя | цель | аргументы | ключевые аргументы |
|---|---|---|---|---|
placeholder | x | x | () | {} |
get_attr | linear_weight | linear.weight | () | {} |
call_function | add_1 | <встроенная функция add> | (x, linear_weight) | {} |
call_module | linear_1 | linear | (add_1,) | {} |
call_method | relu_1 | relu | (linear_1,) | {} |
call_function | sum_1 | <встроенный метод sum …> | (relu_1,) | {‘dim’: -1} |
call_function | topk_1 | <встроенный метод topk …> | (sum_1, 3) | {} |
output | output | output | (topk_1,) | {} |
Мы можем использовать эту информацию для ответа на поставленные выше вопросы.
- Какие входные данные принимаются методом? В FX входные данные метода определяются с помощью специальных
placeholderузлов. В данном случае у нас есть одинplaceholderузел сtargetзначениемx, что означает, что у нас есть один аргумент (без 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)
Для простых преобразований, состоящих только из подстановок, можно также использовать подстановщик подграфов.
Переписывание подграфов с помощью replace_pattern()
FX также предоставляет ещё один уровень автоматизации поверх прямого изменения графа. API replace_pattern() по существу является инструментом «найти/заменить» для редактирования Graph. Он позволяет указать pattern и replacement функцию, и он проследит по этим функциям, найдёт экземпляры группы операций в pattern графе и заменит эти экземпляры копиями replacement графа. Это может значительно автоматизировать утомительный код изменения графа, который становится сложным по мере усложнения преобразований.
Примеры изменения графа
Замена/Переотслеживание
Другой способ манипулирования 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-арными (например, операция сложения является бинарной), так как вызов 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, на который были применены любые правила преобразования, установленные как переопределённые методы.
Примеры шаблона интерпретатора
Отладка
Введение
Часто в процессе написания преобразований наш код не будет работать корректно. В этом случае нам может потребоваться выполнить отладку. Ключ состоит в работе «назад»: сначала проверьте результаты вызова сгенерированного модуля, чтобы доказать или опровергнуть правильность. Затем, просмотрите и отладьте сгенерированный код. Затем отладьте процесс преобразований, которые привели к сгенерированному коду.
Если вы не знакомы с отладчиками, см. дополнительный раздел Доступные отладчики.
Проверка корректности модулей
Поскольку выход большинства модулей глубокого обучения состоит из чисел с плавающей точкой 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 : [#users=1] = placeholder[target=x]
%y : [#users=1] = placeholder[target=y]
%add : [#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» “Python Debugging With Pdb”.
IDE, такие как PyCharm или VSCode, обычно имеют встроенный отладчик. В вашей IDE вы можете выбрать либо a) использовать pdb, открыв окно терминала в вашей IDE (например, View → Terminal в VSCode), или b) использовать встроенный отладчик (обычно графический интерфейс вокруг 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 можно рассматривать как гиперпараметр, и трассировки разных экземпляров MyModule с различными значениями этого параметра имеют различный код. Это допустимый шаблон, который поддерживается символичным прослеживанием.
Многие примеры динамического управления потоком являются семантически статическим управлением потоком. Эти примеры могут быть адаптированы для поддержки символьного прослеживания путем удаления зависимостей данных от входных значений, например, путем перемещения значений в 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, часто аргумент обучения передается как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_allclose(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(), флаг обучения инкапсулируется и — из-за сохранения модели объекта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_allclose(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- Parameters:
-
- root (Union[torch.nn.Module, Callable]) – Модуль или функция, которая будет прослеживаться и преобразовываться в представление графа.
- concrete_args (Optional[Dict[str, any]]) – Входные данные для частичной специализации
- Returns:
-
модуль, созданный из записанных операций из
root. - Return type:
Примечание
Обратная совместимость для этого 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, а не прослеживаются через них.
- Parameters:
-
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()[source] -
Возвращает код Python, сгенерированный для текущего GraphModule и его дочерних GraphModules
Предупреждение
Этот 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 : [#users=1] = self.linear.weight %add_1 : [#users=1] = call_function[target=operator.add](args = (%x, %linear_weight), kwargs = {}) %linear_1 : [#users=1] = call_module[target=linear](args = (%add_1,), kwargs = {}) %relu_1 : [#users=1] = call_method[target=relu](args = (%linear_1,), kwargs = {}) %sum_1 : [#users=1] = call_function[target=torch.sum](args = (%relu_1,), kwargs = {dim: -1}) %topk_1 : [#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. Узел вызова представляет вызов 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. Узел вызова метода представляет вызов заданного метода для 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. Узел вызова модуля представляет вызов функции 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().- Parameters:
-
-
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’. Семантика этих операторных кодов описана в
- Returns:
-
Новый созданный и вставленный узел.
- Return type:
Примечание
Обратная совместимость для этого API гарантирована.
-
eliminate_dead_code()[source] -
Удалить весь неиспользуемый код из графа на основе количества пользователей каждого узла и наличия у узлов побочных эффектов. Граф должен быть топологически отсортирован перед вызовом.
- Returns:
-
Изменился ли граф в результате прохода.
- Return type:
Пример:
Перед удалением неиспользуемого кода узел
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.- Parameters:
-
to_erase (Node) –
Nodeдля удаления изGraph.
Примечание
Обратная совместимость для этого API гарантирована.
-
get_attr(qualified_name, type_expr=None)[source] -
Вставить узел
get_attrв Graph. Узелget_attrNodeпредставляет получение атрибута изModuleиерархии.- Parameters:
-
-
qualified_name (str) – полное имя атрибута для извлечения. Например, если отслеживаемый модуль имеет подмодуль с именем
foo, у которого есть подмодуль с именемbar, у которого есть атрибут с именемbaz, квалифицированное имяfoo.bar.bazдолжно быть передано какqualified_name. - type_expr (Optional[Any]) – необязательная аннотация типа, представляющая Python-тип, который будет иметь вывод этого узла.
-
qualified_name (str) – полное имя атрибута для извлечения. Например, если отслеживаемый модуль имеет подмодуль с именем
- Returns:
-
Новый созданный и вставленный узел
get_attr. - Return type:
Примечание
Те же правила вставки и выражения типа применяются для этого метода, что и для
Graph.create_node.Примечание
Обратная совместимость для этого API гарантирована.
-
graph_copy(g, val_map, return_output_node=False)[source] -
Копирует все узлы из данного графа в
self.- Parameters:
- Returns:
-
Значение в
self, которое теперь эквивалентно выходному значению вg, если уgбыл узелoutput.Noneв противном случае. - Return type:
-
Optional[Union[Tuple[Any, …], List[Any], Dict[str, Any], slice, Node, str, int, float, bool, complex, dtype, Tensor, device, memory_format, layout]]
Примечание
Обратная совместимость для этого 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])- Параметры:
-
-
node (Узел) – Узел, который нужно скопировать в
self. -
arg_transform (Callable[[Узел], Argument]) – Функция, которая преобразует
Nodeаргументы в узлеargsиkwargsв эквивалентные аргументы вself. В простейшем случае это должно извлечь значение из таблицы, сопоставляющей узлы в исходном графе сself.
-
node (Узел) – Узел, который нужно скопировать в
- Тип возвращаемого значения:
Примечание
Совместимость с предыдущими версиями для данного 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 гарантирована.
-
property owning_module -
Возвращает модуль, владеющий этим
GraphModule, если такой существует.Noneесли модуль-владелец отсутствует или если модулей-владельцев несколько.
-
placeholder(name, type_expr=None, default_value)[source] -
Вставка узла
placeholderв граф. Узелplaceholderпредставляет вход функции.- Параметры:
-
-
name (строка) – Имя входного значения. Соответствует имени позиционного аргумента функции, которую представляет данный
Graph. - type_expr (Optional[Any]) – необязательная аннотация типа, представляющая тип Python, который будет иметь результат этого узла. Это необходимо в некоторых случаях для правильной генерации кода (например, при последующем использовании функции в TorchScript).
-
default_value (Any) – Значение по умолчанию для этого аргумента функции. ПРИМЕЧАНИЕ: чтобы разрешить
Noneв качестве значения по умолчанию,inspect.Signature.emptyдолжен быть передан в качестве этого аргумента, чтобы указать, что у параметра _нет_ значения по умолчанию.
-
name (строка) – Имя входного значения. Соответствует имени позиционного аргумента функции, которую представляет данный
- Тип возвращаемого значения:
Примечание
Те же правила вставки и выражений типа применяются для данного метода, что и для
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) – Имя корневого модуля, в котором нужно искать квалифицированные имена целевых объектов. Обычно это ‘self’.
- Возвращает:
-
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, Node, str, int, float, bool, complex, dtype, Tensor, device, memory_format, layout]], ...] -
Кортеж аргументов этого
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[Список[str]]) – Список, который будет хранить отформатированные строки, представляющие заполнители в сгенерированной функции
forward. Использование только внутри. -
maybe_return_typename (Optional[Список[str]]) – Список с одним элементом, который будет хранить отформатированную строку, представляющую выходную функцию
forward. Использование только внутри.
-
placeholder_names (Optional[Список[str]]) – Список, который будет хранить отформатированные строки, представляющие заполнители в сгенерированной функции
- Возвращает:
-
-
If 1) we’re using format_node as an internal helper -
в методе
__str__Graph, и 2)selfэто узел-заполнитель, вернётсяNone. В противном случае вернёт описательную строковую запись текущего узла.
-
- Тип возвращаемого значения:
Примечание
Обратная совместимость этого API гарантирована.
-
is_impure()[source] -
Возвращает, является ли эта операция нечистой, т.е. если её операция это заполнитель или вывод, или если call_function или call_module, который является нечистым.
- Возвращает:
-
Если операция нечистая или нет.
- Тип возвращаемого значения:
Предупреждение
Этот API экспериментальный и НЕ обратной совместимости.
-
property kwargs: Dict[str, Optional[Union[Tuple[Any, ...], List[Any], Dict[str, Any], slice, Node, str, int, float, bool, complex, dtype, Tensor, device, memory_format, layout]]] -
Словарь ключевых аргументов этого
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будут сопоставлены с сигнатурой модуля/функции и возвращать исключительно значения 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]]) – Словарь типов аргументов для kwargs.
- normalize_to_only_use_kwargs (bool) – Нужно ли нормализовать для использования только kwargs.
- Возвращаемое значение:
-
Возвращает кортеж 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>>)[source] -
Заменить все использования
selfв графе на узелreplace_with.- Параметры:
-
-
replace_with (Узел) – Узел, на который нужно заменить все использования
self. - delete_user_cb (Callable) – Обратный вызов, который вызывается для определения, нужно ли удалять данный пользователь узла self.
-
replace_with (Узел) – Узел, на который нужно заменить все использования
- Возвращаемое значение:
-
Список узлов, на которых было произведено это изменение.
- Тип возвращаемого значения:
Примечание
Обратная совместимость для этого API гарантирована.
-
replace_input_with(old_input, new_input)[source] -
Проходим по входным узлам
self, и заменяем все экземплярыold_inputнаnew_input.- Параметры:
Примечание
Обратная совместимость для этого API гарантирована.
-
property stack_trace: Optional[str] -
Возвращает стек вызовов Python, который был записан во время трассировки, если таковой имеется. Это свойство обычно заполняется
Tracer.create_proxy. Чтобы записывать стеки вызовов во время трассировки для отладки, установитеrecord_stack_traces = Trueна экземпляреTracer.
-
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 (Вызываемый объект) – Метод forward()
Moduleдля вызова - args (Кортеж) – args вызова модуля
- kwargs (Словарь) – kwargs вызова модуля
- Возвращаемое значение:
-
Значение, возвращаемое вызовом модуля. В случае генерации узла
call_module, это значениеProxy. В противном случае, это значение, возвращаемое вызовомModule. - Тип возвращаемого значения:
Примечание
Обратная совместимость этого API гарантирована.
-
create_arg(a)[source] -
Метод для определения поведения трассирования при подготовке значений для использования в качестве аргументов узлов в
Graph.По умолчанию, поведение включает:
- Итерацию по коллекциям (например, кортежам, спискам, словарям) и рекурсивный вызов
create_argsдля элементов. - Для объекта Proxy возвращается ссылка на подлежащий IR
Node. -
Для объекта Tensor, не являющегося Proxy, генерируется IR для различных случаев:
- Для параметра Parameter генерируется узел
get_attr, ссылающийся на этот параметр. - Для Tensor, не являющегося параметром Parameter, Tensor сохраняется в специальном атрибуте, ссылающемся на этот атрибут.
- Для параметра Parameter генерируется узел
Этот метод можно переопределить для поддержки дополнительных типов.
- Параметры:
-
a (Любой тип) – Значение, которое должно быть сгенерировано как
ArgumentвGraph. - Возвращаемое значение:
-
Значение
aпреобразованное в соответствующееArgument. - Тип возвращаемого значения:
-
Optional[Union[Tuple[Любой тип, …], Список[Любой тип], Словарь[str, Любой тип], slice, Узел, str, int, float, bool, complex, Тип данных, Тензор, Устройство, Формат памяти, Макет]]
Примечание
Обратная совместимость этого API гарантирована.
- Итерацию по коллекциям (например, кортежам, спискам, словарям) и рекурсивный вызов
-
create_args_for_root(root_fn, is_module, concrete_args=None)[source] -
Создаёт узлы
placeholderсоответствующие сигнатуре модуляroot. Этот метод инспектирует сигнатуру корневого модуля и генерирует узлы соответственно, поддерживая*argsи**kwargs.Предупреждение
Этот API экспериментальный и НЕ обратной совместимости.
-
create_node(kind, target, args, kwargs, name=None, type_expr=None) -
Вставляет узел графа, заданный целевым значением, аргументами, ключевыми аргументами и именем.
Этот метод можно переопределить для выполнения дополнительных проверок, валидации или модификации значений, используемых при создании узла. Например, можно запретить запись операций in-place.
Примечание
Обратная совместимость этого API гарантирована.
- Тип возвращаемого значения:
-
create_proxy(kind, target, args, kwargs, name=None, type_expr=None, proxy_factory_fn=None) -
Создает узел Node из заданных аргументов, а затем возвращает узел Node, заключенный в объект Proxy.
Если kind = ‘placeholder’, то создается узел Node, представляющий параметр функции. Если нам нужно закодировать параметр по умолчанию, мы используем кортеж
args.argsв противном случае пуст для узловplaceholder.Примечание
Обратная совместимость этого API гарантирована.
-
getattr(attr, attr_val, parameter_proxy_cache)[source] -
Метод, определяющий поведение данного
Tracerпри вызове getattr для экземпляраnn.Module.По умолчанию, возвращает прокси-значение для атрибута. Также сохраняет прокси-значение в
parameter_proxy_cache, для повторного использования в будущих вызовах, вместо создания нового.Этот метод можно переопределить, например, для того, чтобы не возвращать прокси при запросе параметров.
- Параметры:
- Возвращаемое значение:
-
Возвращаемое значение вызова getattr.
Предупреждение
Этот API экспериментальный и НЕ обратной совместимости.
-
-
is_leaf_module(m, module_qualified_name)[source] -
Метод для указания, является ли данный
nn.Module«листовым» модулем.Листовые модули являются атомарными единицами, которые появляются в IR, на которые ссылаются вызовы
call_module. По умолчанию модули в пространстве имен стандартной библиотеки PyTorch (torch.nn) являются листовыми модулями. Все остальные модули отслеживаются и записываются их составляющие операции, если явно не указано иное с помощью этого параметра.- Параметры:
-
- m (Модуль) – Модуль, о котором запрашивается информация
-
module_qualified_name (строка) – Путь к корню этого модуля. Например, если у вас есть иерархия модулей, где подмодуль
fooсодержит подмодульbar, который содержит подмодульbaz, этот модуль будет отображаться с квалифицированным именемfoo.bar.bazздесь.
- Тип возвращаемого значения:
Примечание
Обратная совместимость для этого API гарантирована.
-
iter(obj) -
- Вызывается, когда объект-прокси итерируется, например,
-
когда используется в управляющей последовательности. Обычно мы не знаем, что делать, потому что не знаем значение прокси, но пользовательский отслеживатель может добавить больше информации к узлу графа с помощью create_node и может выбрать возвращение итератора.
Примечание
Обратная совместимость для этого API гарантирована.
- Тип возвращаемого значения:
-
keys(obj) -
- Вызывается, когда для объекта-прокси вызывается метод keys().
-
Это происходит, когда вызывается ** для прокси. Это должно вернуть итератор, если ** должен работать в вашем пользовательском отслеживателе.
Примечание
Обратная совместимость для этого API гарантирована.
- Тип возвращаемого значения:
-
path_of_module(mod)[source] -
Вспомогательный метод для поиска квалифицированного имени
modв иерархии модулейroot. Например, еслиrootимеет подмодуль с именемfoo, который имеет подмодуль с именемbar, передачаbarв эту функцию вернет строку «foo.bar».- Параметры:
-
mod (строка) –
Moduleдля получения квалифицированного имени. - Тип возвращаемого значения:
Примечание
Обратная совместимость для этого 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[Модуль, Вызов]) – Либо
Module, либо функция, которую нужно отследить. Обратная совместимость для этого параметра гарантирована. - concrete_args (Optional[Dict[строка, любой]]) – Конкретные аргументы, которые не должны обрабатываться как прокси. Этот параметр экспериментальный, и его обратная совместимость НЕ гарантирована.
-
root (Union[Модуль, Вызов]) – Либо
- Возвращает:
-
Graph, представляющее семантику переданного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_allclose(result, torch.neg(input).sigmoid())- Параметры:
-
- module (GraphModule) – Модуль, который нужно выполнить
-
garbage_collect_values (bool) – Удалять ли значения после их последнего использования в рамках выполнения Модуля. Это гарантирует оптимальное использование памяти во время выполнения. Это можно отключить, например, для проверки всех промежуточных значений в выполнении, посмотрев на атрибут
Interpreter.env.
Примечание
Обратная совместимость для этого 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[Кортеж[Любой, …], Список[Любой], Словарь[строка, Любой], срезы, Узел, строка, целое, вещественное, булево, комплексное, тип_данных, Тензор, устройство, форма_памяти, макет]]
Примечание
Обратная совместимость для этого 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 (булево) – Если 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_allclose(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 (Callable) — Подграф для сопоставления в
gmдля замены -
replacement (Callable) — Подграф для замены
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, в то время как последним узлом является то, что возвращает функция.)Важно отметить, что параметры Callable
patternдолжны использоваться в самом Callable, а параметры Callablereplacementдолжны соответствовать шаблону. Первое правило объясняет, почему в приведенном коде функция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/1.13/fx.html