torch.fx
Создано: 15 декабря 2020 г. | Последнее обновление: 8 мая 2026 г.
Обзор
FX — это набор инструментов, позволяющий разработчикам преобразовывать экземпляры nn.Module. FX состоит из трех основных компонентов: символического трассировщика, промежуточного представления и генерации кода Python. Пример работы этих компонентов:
import torch
# Simple module for demonstration
class MyModule(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.param = torch.nn.Parameter(torch.rand(3, 4))
self.linear = torch.nn.Linear(4, 5)
def forward(self, x):
return self.linear(x + self.param).clamp(min=0.0, max=1.0)
module = MyModule()
from torch.fx import symbolic_trace
# Symbolic tracing frontend - captures the semantics of the module
symbolic_traced: torch.fx.GraphModule = symbolic_trace(module)
# High-level intermediate representation (IR) - Graph representation
print(symbolic_traced.graph)
"""
graph():
%x : [num_users=1] = placeholder[target=x]
%param : [num_users=1] = get_attr[target=param]
%add : [num_users=1] = call_function[target=operator.add](args = (%x, %param), kwargs = {})
%linear : [num_users=1] = call_module[target=linear](args = (%add,), kwargs = {})
%clamp : [num_users=1] = call_method[target=clamp](args = (%linear,), kwargs = {min: 0.0, max: 1.0})
return clamp
"""
# Code generation - valid Python code
print(symbolic_traced.code)
"""
def forward(self, x):
param = self.param
add = x + param; x = param = None
linear = self.linear(add); add = None
clamp = linear.clamp(min = 0.0, max = 1.0); linear = None
return clamp
"""
Символический трассировщик выполняет «символическое исполнение» кода Python. Он передает через код фиктивные значения, называемые прокси-объектами (Proxy). Операции над этими прокси-объектами записываются. Подробнее о символической трассировке см. в документации по symbolic_trace() и Tracer.
Промежуточное представление — это контейнер для операций, записанных во время символической трассировки. Оно состоит из списка узлов (Node), представляющих входные данные функции, места вызова (функций, методов или экземпляров torch.nn.Module) и возвращаемые значения. Подробнее о IR см. в документации по Graph. IR — это формат, к которому применяются преобразования.
Генерация кода Python превращает FX в инструментарий для преобразований Python в Python (или модулей в модули). Для каждого графа IR можно создать корректный код 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. Возвращаемый преобразованием FX torch.nn.Module следует считать идентичным обычному torch.nn.Module: его можно передать другому преобразованию FX или запустить. Если на вход и выход преобразования FX подавать torch.nn.Module, преобразования можно будет комбинировать.
Примечание
Также можно изменить существующий GraphModule, а не создавать новый, например, так:
import torch
import torch.fx
def transform(m : nn.Module) -> nn.Module:
gm : torch.fx.GraphModule = torch.fx.symbolic_trace(m)
# Modify gm.graph
# <...>
# Recompile the forward() method of `gm` from its Graph
gm.recompile()
return gm
Обратите внимание: необходимо вызвать GraphModule.recompile(), чтобы сгенерированный метод forward() для GraphModule соответствовал измененному Graph.
Если вы передали torch.nn.Module, трассированный в Graph, то теперь можно выбрать один из двух основных подходов к созданию нового Graph.
Краткое введение в графы
Полное описание семантики графов приведено в документации по Graph, а здесь мы рассмотрим основные понятия. Graph — это структура данных, представляющая метод в GraphModule. Для этого нужно указать:
- Каковы входные данные метода?
- Какие операции выполняются внутри метода?
- Каково выходное значение (то есть возвращаемое значение) метода?
Все три понятия представлены экземплярами Node. Рассмотрим короткий пример:
import torch
import torch.fx
class MyModule(torch.nn.Module):
def __init__(self):
super().__init__()
self.param = torch.nn.Parameter(torch.rand(3, 4))
self.linear = torch.nn.Linear(4, 5)
def forward(self, x):
return torch.topk(torch.sum(
self.linear(x + self.linear.weight).relu(), dim=-1), 3)
m = MyModule()
gm = torch.fx.symbolic_trace(m)
gm.graph.print_tabular()
Здесь мы определяем модуль MyModule для демонстрации, создаем его экземпляр, выполняем символическую трассировку, а затем вызываем метод Graph.print_tabular(), чтобы вывести таблицу с узлами этого Graph:
opcode | name | target | args | kwargs |
|---|---|---|---|---|
placeholder | x | x | () | {} |
get_attr | linear_weight | linear.weight | () | {} |
call_function | add_1 | (x, linear_weight) | {} | |
call_module | linear_1 | linear | (add_1,) | {} |
call_method | relu_1 | relu | (linear_1,) | {} |
call_function | sum_1 | <built-in method sum …> | (relu_1,) | {‘dim’: -1} |
call_function | topk_1 | <built-in method topk …> | (sum_1, 3) | {} |
output | output | output | (topk_1,) | {} |
Эти сведения помогут ответить на поставленные выше вопросы.
- Каковы входные данные метода? В FX входные данные метода задаются специальными узлами
placeholder. В этом случае у нас есть один узелplaceholderс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\s. Он позволяет задать функцию pattern и функцию replacement, затем трассирует эти функции, находит в графе pattern группы операций и заменяет их копиями графа replacement. Это позволяет значительно автоматизировать утомительный код для манипуляций с графом, который может стать громоздким по мере усложнения преобразований.
Примеры манипуляций с графом
Прокси и повторная трассировка
Еще один способ манипулировать Graph\s — повторно использовать механизм 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\s дает возможность задавать правила переписывания в виде обычного кода Python. Для преобразований, требующих большого количества правил (например, vmap или grad), это зачастую повышает читаемость и удобство сопровождения. Обратите внимание: при вызове Proxy мы также передали трассировщик, указывающий на базовую переменную graph. Это нужно для того, чтобы в случае n-арных операций в графе (например, add — бинарный оператор) вызов Proxy не создавал несколько экземпляров трассировщика графа, что может приводить к непредвиденным ошибкам во время выполнения. Мы рекомендуем использовать Proxy таким способом, особенно если нельзя с уверенностью считать, что базовые операторы унарны.
Пример использования Proxy\s для манипуляций с Graph можно найти здесь.
Шаблон «Интерпретатор»
Полезный шаблон организации кода в FX — перебор и выполнение всех Node\s в Graph. Это можно использовать для разных целей, в том числе для анализа значений, проходящих через граф во время выполнения, или для преобразования кода путем повторной трассировки с помощью Proxy\s. Например, предположим, что мы хотим запустить 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 nonexistent 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
"""
Здесь мы попытались проверить равенство значений двух моделей глубокого обучения с помощью оператора равенства ==. Однако такое сравнение некорректно: оператор возвращает тензор, а не bool; кроме того, при сравнении чисел с плавающей точкой нужно учитывать погрешность (или эпсилон), поскольку операции с плавающей точкой некоммутативны (подробнее см. здесь). Вместо этого можно использовать torch.allclose(), которая выполняет приблизительное сравнение с учетом относительного и абсолютного порогов допуска:
assert torch.allclose(resnet18(input_image), transformed_resnet18(input_image))
Это первый инструмент, который поможет проверить, работают ли преобразованные модули так, как мы ожидаем, по сравнению с эталонной реализацией.
Отладка сгенерированного кода
Поскольку FX генерирует функцию forward() для GraphModule\s, традиционные методы отладки, такие как инструкции print или pdb, не так просто применить. К счастью, для отладки сгенерированного кода можно использовать несколько приемов.
Использование pdb
Вызовите pdb, чтобы перейти к выполнению программы в отладчике. Хотя код, представляющий Graph, отсутствует в каком-либо исходном файле, при вызове forward pass можно вручную перейти к нему с помощью 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 pass в свой код и изучить его там.
# 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 в папку. Хотя копирования forward pass в код часто достаточно, как описано в разделе Вывод сгенерированного кода, иногда с помощью to_folder удобнее изучать модули и параметры.
m = symbolic_trace(M())
m.to_folder("foo", "Bar")
from foo import Bar
y = Bar()
После выполнения примера выше можно изучить код в foo/module.py и при необходимости изменить его (например, добавить инструкции print или использовать pdb), чтобы отладить сгенерированный код.
Отладка преобразования
Теперь, когда мы установили, что преобразование создает некорректный код, пора отладить само преобразование. Сначала изучим раздел документации Ограничения символической трассировки. Убедившись, что трассировка работает как ожидалось, нужно выяснить, что пошло не так при преобразовании GraphModule. Возможно, ответ найдется в разделе Написание преобразований; если нет, есть несколько способов изучить трассированный модуль:
# Sample Module
class M(torch.nn.Module):
def forward(self, x, y):
return x + y
# Create an instance of `M`
m = M()
# Symbolically trace an instance of `M` (returns a GraphModule). In
# this example, we'll only be discussing how to inspect a
# GraphModule, so we aren't showing any sample transforms for the
# sake of brevity.
traced = symbolic_trace(m)
# Print the code produced by tracing the module.
print(traced)
# The generated `forward` function is:
"""
def forward(self, x, y):
add = x + y; x = y = None
return add
"""
# Print the internal Graph.
print(traced.graph)
# This print-out returns:
"""
graph():
%x : [num_users=1] = placeholder[target=x]
%y : [num_users=1] = placeholder[target=y]
%add : [num_users=1] = call_function[target=operator.add](args = (%x, %y), kwargs = {})
return add
"""
# Print a tabular representation of the internal Graph.
traced.graph.print_tabular()
# This gives us:
"""
opcode name target args kwargs
------------- ------ ----------------------- ------ --------
placeholder x x () {}
placeholder y y () {}
call_function add <built-in function add> (x, y) {}
output output output (add,) {}
"""
С помощью приведенных выше вспомогательных функций можно сравнить трассированный модуль до и после применения преобразований. Иногда для обнаружения ошибки достаточно простого визуального сравнения. Если причина проблемы все еще неясна, следующим шагом может стать отладчик, например pdb.
Рассмотрим следующий код на основе приведенного выше примера:
# Sample user-defined function
def transform_graph(module: torch.nn.Module, tracer_class : type = fx.Tracer) -> torch.nn.Module:
# Get the Graph from our traced Module
g = tracer_class().trace(module)
"""
Transformations on `g` go here
"""
return fx.GraphModule(module, g)
# Transform the Graph
transformed = transform_graph(traced)
# Print the new code after our transforms. Check to see if it was
# what we expected
print(transformed)
Предположим, что вызов print(traced) показал ошибку в наших преобразованиях. Чтобы выяснить причину с помощью отладчика, запустим сеанс pdb. Чтобы проследить за ходом преобразования, установим точку останова на transform_graph(traced), а затем нажмем s, чтобы «войти» в вызов transform_graph(traced).
Также может помочь изменение метода print_tabular, чтобы выводить различные атрибуты узлов графа. (Например, могут понадобиться input_nodes и users узла.)
Доступные отладчики
Самый распространенный отладчик Python — pdb. Запустить программу в «режиме отладки» с помощью pdb можно, введя в командной строке python -m pdb FILENAME.py, где FILENAME — имя файла, который нужно отладить. После этого можно использовать pdb команды отладчика, чтобы пошагово выполнять программу. Обычно при запуске pdb устанавливают точку останова (b LINE-NUMBER), а затем вызывают c, чтобы выполнить программу до этой точки. Так не придется проходить каждую строку кода с помощью s или n, чтобы добраться до нужного участка. Другой вариант — добавить import pdb; pdb.set_trace() перед строкой, на которой нужно остановиться. Если добавить pdb.set_trace(), программа автоматически запустится в режиме отладки. (Иными словами, в командной строке можно ввести python FILENAME.py вместо python -m pdb FILENAME.py.) После запуска файла в режиме отладки можно пошагово выполнять код и изучать внутреннее состояние программы с помощью определенных команд. В интернете есть множество отличных руководств по pdb, в том числе статья RealPython «Отладка Python с помощью Pdb».
В IDE, таких как PyCharm или VSCode, обычно встроен отладчик. В своей IDE можно либо а) использовать pdb, открыв окно терминала в IDE (например, «Вид» → «Терминал» в VSCode), либо б) воспользоваться встроенным отладчиком (обычно это графическая оболочка для pdb).
Ограничения символической трассировки
FX использует систему символической трассировки (также известной как символическое выполнение) для фиксации семантики программ в форме, пригодной для преобразования и анализа. Эта система является трассировкой, поскольку она выполняет программу (точнее, torch.nn.Module или функцию), чтобы записывать операции. Она является символической, поскольку данные, проходящие через программу во время выполнения, — это не реальные данные, а символы (в терминологии FX — Proxy).
Хотя символическая трассировка работает с большинством кода нейронных сетей, у неё есть некоторые ограничения.
Динамический поток управления
Основное ограничение символической трассировки заключается в том, что в настоящее время она не поддерживает динамический поток управления. То есть циклы или инструкции 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и подмодулей- При использовании функциональных API, таких как
torch.nn.functional.dropout, аргумент training часто передаётся какself.training. Во время трассировки FX это значение, скорее всего, будет зафиксировано как константа.
import torch import torch.fx class DropoutRepro(torch.nn.Module): def forward(self, x): return torch.nn.functional.dropout(x, training=self.training) traced = torch.fx.symbolic_trace(DropoutRepro()) print(traced.code) """ def forward(self, x): dropout = torch.nn.functional.dropout(x, p = 0.5, training = True, inplace = False); x = None return dropout """ traced.eval() x = torch.randn(5, 3) torch.testing.assert_close(traced(x), x) """ AssertionError: Tensor-likes are not close! Mismatched elements: 15 / 15 (100.0%) Greatest absolute difference: 1.6207983493804932 at index (0, 2) (up to 1e-05 allowed) Greatest relative difference: 1.0 at index (0, 0) (up to 0.0001 allowed) """- Однако при использовании стандартного подмодуля
nn.Dropout()флаг training инкапсулирован и — благодаря сохранению объектной моделиnn.Module— может изменяться.
class DropoutRepro2(torch.nn.Module): def __init__(self): super().__init__() self.drop = torch.nn.Dropout() def forward(self, x): return self.drop(x) traced = torch.fx.symbolic_trace(DropoutRepro2()) print(traced.code) """ def forward(self, x): drop = self.drop(x); x = None return drop """ traced.eval() x = torch.randn(5, 3) torch.testing.assert_close(traced(x), x) - При использовании функциональных API, таких как
- Из-за этого различия рассмотрите возможность пометить модули, динамически взаимодействующие с флагом
training, как листовые модули.
Справочник API
-
torch.fx.symbolic_trace(root, concrete_args=None)[исходный код] -
API символической трассировки
Для экземпляра
nn.Moduleили функцииrootэта функция вернётGraphModule, созданный путём записи операций, обнаруженных при трассировкеroot.concrete_argsпозволяет частично специализировать функцию, например, чтобы устранить поток управления или структуры данных.Например:
def f(a, b): if b == True: return a else: return a * 2FX обычно не может трассировать этот код из-за наличия потока управления. Однако с помощью
concrete_argsможно специализировать трассировку по значениюb:f = fx.symbolic_trace(f, concrete_args={"b": False}) assert f(3, False) == 6Обратите внимание, что, хотя вы по-прежнему можете передавать разные значения
b, они будут игнорироваться.С помощью
concrete_argsтакже можно устранить обработку структур данных в функции. Для этого будут использоваться pytrees, чтобы развернуть входные данные. Чтобы избежать чрезмерной специализации, передавайтеfx.PHдля значений, которые не следует специализировать. Например:def f(x): out = 0 for v in x.values(): out += v return out f = fx.symbolic_trace( f, concrete_args={"x": {"a": fx.PH, "b": fx.PH, "c": fx.PH}} ) assert f({"a": 1, "b": 2, "c": 4}) == 7- Параметры:
-
- root (Union[torch.nn.Module, Callable]) – Модуль или функция для трассировки и преобразования в представление Graph.
- concrete_args (Optional[Dict[str, any]]) – Входные данные для частичной специализации
- Возвращает:
-
Модуль, созданный на основе операций, записанных из
root. - Тип возвращаемого значения:
Примечание
Обратная совместимость этого API гарантируется.
-
torch.fx.wrap(fn_or_name: _F) → _F[исходный код] - torch.fx.wrap(fn_or_name:str) str
-
Эту функцию можно вызвать на уровне модуля, чтобы зарегистрировать fn_or_name как «листовую функцию». «Листовая функция» будет сохранена в трассировке FX как узел CallFunction, а не трассироваться изнутри:
# foo/bar/baz.py def my_custom_function(x, y): return x * x + y * y torch.fx.wrap("my_custom_function") def fn_to_be_traced(x, y): # When symbolic tracing, the below call to my_custom_function will be inserted into # the graph rather than tracing it. return my_custom_function(x, y)Эту функцию также можно использовать в качестве декоратора:
# foo/bar/baz.py @torch.fx.wrap def my_custom_function(x, y): return x * x + y * yФункцию-обёртку можно считать «листовой функцией», аналогичной понятию «листовых модулей»: такие функции остаются в трассировке FX в виде вызовов, а не трассируются изнутри.
- Параметры:
-
fn_or_name (Union[str, Callable]) – Функция или имя глобальной функции, которую нужно добавить в граф при её вызове
Примечание
Обратная совместимость этого API гарантируется.
-
class torch.fx.GraphModule(*args, **kwargs)[исходный код] -
GraphModule — это nn.Module, созданный на основе fx.Graph. GraphModule имеет атрибут
graph, а также атрибутыcodeиforward, созданные на основе этогоgraph.Предупреждение
При повторном присваивании
graphатрибутыcodeиforwardбудут автоматически созданы заново. Однако если изменить содержимоеgraph, не переназначая сам атрибутgraph, необходимо вызватьrecompile(), чтобы обновить сгенерированный код.Примечание
Обратная совместимость этого API гарантируется.
- Тип возвращаемого значения:
-
Self
-
__init__(root, graph, class_name='GraphModule')[исходный код] -
Создать GraphModule.
- Параметры:
-
-
root (Union[torch.nn.Module, Dict[str, Any]) –
rootможет быть экземпляром nn.Module или Dict, сопоставляющим строки с атрибутами любого типа. Еслиroot— это Module, все ссылки на объекты на основе Module (по квалифицированному имени) в полеtargetузлов Graph будут скопированы из соответствующих мест иерархии модулейrootв иерархию модулей GraphModule. Еслиroot— это dict, квалифицированное имя, указанное вtargetузла, будет напрямую искаться среди ключей словаря. Объект, связанный с этим ключом в Dict, будет скопирован в соответствующее место иерархии модулей 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)[исходный код] -
Добавляет указанный подмодуль в
self.Если ещё не существуют промежуточные модули, являющиеся частью пути
target, для них создаются пустые Module.- Параметры:
- Возвращает:
-
- Удалось ли вставить подмодуль. Чтобы
-
этот метод вернул True, каждый объект в цепочке, обозначенной
target, должен либо а) ещё не существовать, либо б) ссылаться наnn.Module(а не на параметр или другой атрибут)
- Тип возвращаемого значения:
Примечание
Обратная совместимость этого API гарантируется.
-
property code: str -
Возвращает код Python, сгенерированный на основе
Graph, лежащего в основе этогоGraphModule.
-
delete_all_unused_submodules()[исходный код] -
Удаляет все неиспользуемые подмодули из
self.Модуль считается «используемым», если выполняется хотя бы одно из условий: 1. У него есть используемые дочерние модули. 2. Его forward вызывается напрямую через узел
call_module. 3. У него есть атрибут, не являющийся Module, который используется из узлаget_attr.Этот метод можно вызвать для очистки
nn.Moduleбез необходимости вручную вызыватьdelete_submoduleдля каждого неиспользуемого подмодуля.Примечание
Обратная совместимость этого API гарантируется.
-
delete_submodule(target)[исходный код] -
Удаляет указанный подмодуль из
self.Модуль не будет удалён, если
targetне является допустимой целью.- Параметры:
-
target (str) – Полное строковое имя нового подмодуля (пример задания полного строкового имени см. в
nn.Module.get_submodule) - Возвращает:
-
- Указывает, ссылается ли строка target на
-
подмодуль, который нужно удалить. Возвращаемое значение
Falseозначает, чтоtargetне является допустимой ссылкой на подмодуль.
- Тип возвращаемого значения:
Примечание
Обратная совместимость этого API гарантируется.
-
property graph: Graph -
Возвращает
Graph, лежащий в основе этогоGraphModule
-
print_readable(print_output=True, include_stride=False, include_device=False, colored=False, *, fast_sympy_print=False, expanded_def=False, additional_meta=None)[исходный код] -
Возвращает код Python, сгенерированный для текущего GraphModule и его дочерних GraphModule.
- Параметры:
-
additional_meta (list[str] | None) – Необязательный список ключей метаданных, которые нужно включить в вывод. Для каждого ключа в списке, если он есть в node.meta, его значение будет отображено в формате «ключ: значение». Пример:
print_readable(additional_meta=[“seq_nr”]). - Тип возвращаемого значения:
Предупреждение
Этот API является экспериментальным и НЕ гарантирует обратную совместимость.
-
recompile()[исходный код] -
Повторно компилирует этот GraphModule на основе его атрибута
graph. Этот метод следует вызывать после редактирования содержащегося в нёмgraph, иначе сгенерированный код этогоGraphModuleбудет устаревшим.Примечание
Обратная совместимость этого API гарантируется.
- Тип возвращаемого значения:
-
PythonCode
-
to_folder(folder, module_name='FxModule')[исходный код] -
-
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)[исходный код] -
Graph— это основная структура данных, используемая в промежуточном представлении FX. Она состоит из последовательностиNode, каждый из которых представляет место вызова (или другую синтаксическую конструкцию). СписокNode, взятый целиком, образует корректную функцию Python.Например, следующий код
import torch import torch.fx class MyModule(torch.nn.Module): def __init__(self): super().__init__() self.param = torch.nn.Parameter(torch.rand(3, 4)) self.linear = torch.nn.Linear(4, 5) def forward(self, x): return torch.topk( torch.sum(self.linear(x + self.linear.weight).relu(), dim=-1), 3 ) m = MyModule() gm = torch.fx.symbolic_trace(m)создаст следующий граф:
print(gm.graph)
graph(x): %linear_weight : [num_users=1] = self.linear.weight %add_1 : [num_users=1] = call_function[target=operator.add](args = (%x, %linear_weight), kwargs = {}) %linear_1 : [num_users=1] = call_module[target=linear](args = (%add_1,), kwargs = {}) %relu_1 : [num_users=1] = call_method[target=relu](args = (%linear_1,), kwargs = {}) %sum_1 : [num_users=1] = call_function[target=torch.sum](args = (%relu_1,), kwargs = {dim: -1}) %topk_1 : [num_users=1] = call_function[target=torch.topk](args = (%sum_1, 3), kwargs = {}) return topk_1Описание семантики операций, представленных в
Graph, см. в разделеNode.Примечание
Для этого API гарантирована обратная совместимость.
-
__init__(owning_module=None, tracer_cls=None, tracer_extras=None)[исходный код] -
Создать пустой граф.
Примечание
Для этого API гарантирована обратная совместимость.
-
call_function(the_function, args=None, kwargs=None, type_expr=None, name=None)[исходный код] -
Вставить
call_functionNodeвGraph. Узелcall_functionпредставляет вызов вызываемого объекта Python, заданного с помощьюthe_function.- Параметры:
-
-
the_function (Callable[..., Any]) – Вызываемая функция. Это может быть любой оператор PyTorch, функция Python или член пространств имен
builtinsилиoperator. - args (Optional[Tuple[Argument, ...]]) – Позиционные аргументы, передаваемые вызываемой функции.
- kwargs (Optional[Dict[str, Argument]]) – Именованные аргументы, передаваемые вызываемой функции
- type_expr (Optional[Any]) – необязательная аннотация типа, представляющая тип Python, который будет иметь результат этого узла.
- name (Optional[str]) – Имя узла. Если не указано, устанавливается значение None
-
the_function (Callable[..., Any]) – Вызываемая функция. Это может быть любой оператор PyTorch, функция Python или член пространств имен
- Возвращает:
-
Недавно созданный и вставленный узел
call_function. - Тип возвращаемого значения:
Примечание
Для этого метода действуют те же правила точки вставки и выражения типа, что и для
Graph.create_node().Примечание
Для этого API гарантирована обратная совместимость.
-
call_method(method_name, args=None, kwargs=None, type_expr=None)[исходный код] -
Вставить
call_methodNodeвGraph. Узелcall_methodпредставляет вызов заданного метода для элемента с индексом 0 вargs.- Параметры:
-
-
method_name (str) – Имя метода, применяемого к аргументу self. Например, если args[0] — это
Node, представляющийTensor, то для вызоваrelu()для этогоTensorпередайтеreluвmethod_name. -
args (Optional[Tuple[Argument, ...]]) – Позиционные аргументы, передаваемые вызываемому методу. Обратите внимание, что сюда следует включить аргумент
self. - kwargs (Optional[Dict[str, Argument]]) – Именованные аргументы, передаваемые вызываемому методу
- type_expr (Optional[Any]) – необязательная аннотация типа, представляющая тип Python, который будет иметь результат этого узла.
-
method_name (str) – Имя метода, применяемого к аргументу self. Например, если args[0] — это
- Возвращает:
-
Недавно созданный и вставленный узел
call_method. - Тип возвращаемого значения:
Примечание
Для этого метода действуют те же правила точки вставки и выражения типа, что и для
Graph.create_node().Примечание
Для этого API гарантирована обратная совместимость.
-
call_module(module_name, args=None, kwargs=None, type_expr=None)[исходный код] -
Вставить
call_moduleNodeвGraph. Узелcall_moduleпредставляет вызов функции forward() объектаModuleв иерархииModule.- Параметры:
-
-
module_name (str) – Полное квалифицированное имя объекта
Moduleв иерархииModule, который требуется вызвать. Например, если трассированныйModuleсодержит подмодуль с именемfoo, у которого есть подмодуль с именемbar, в качествеmodule_nameдля вызова этого модуля следует передать полное квалифицированное имяfoo.bar. -
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)[исходный код] -
Создать
Nodeи добавить его вGraphв текущей точке вставки. Обратите внимание, что текущую точку вставки можно задать с помощьюGraph.inserting_before()иGraph.inserting_after().- Параметры:
-
-
op (str) – код операции этого узла. Одно из значений: ‘call_function’, ‘call_method’, ‘get_attr’, ‘call_module’, ‘placeholder’ или ‘output’. Семантика этих кодов операций описана в docstring
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’. Семантика этих кодов операций описана в docstring
- Возвращает:
-
Недавно созданный и вставленный узел.
- Тип возвращаемого значения:
Примечание
Для этого API гарантирована обратная совместимость.
-
create_size_node(tensor_node, dim)[исходный код] -
Создать узел FX для
tensor_node.size(dim).Предупреждение
Этот API является экспериментальным и НЕ обеспечивает обратную совместимость.
- Тип возвращаемого значения:
-
create_storage_offset_node(tensor_node)[исходный код] -
Создать узел FX для
tensor_node.storage_offset().Предупреждение
Этот API является экспериментальным и НЕ обеспечивает обратную совместимость.
- Тип возвращаемого значения:
-
create_stride_node(tensor_node, dim)[исходный код] -
Создать узел FX для
tensor_node.stride(dim).Предупреждение
Этот API является экспериментальным и НЕ обеспечивает обратную совместимость.
- Тип возвращаемого значения:
-
eliminate_dead_code(is_impure_node=None)[исходный код] -
Удалить из графа весь неиспользуемый код, учитывая количество пользователей каждого узла и наличие у узлов побочных эффектов. Перед вызовом граф должен быть отсортирован топологически.
- Параметры:
- Возвращает:
-
Был ли граф изменён в результате этого прохода.
- Тип возвращаемого значения:
Пример:
До удаления неиспользуемого кода у
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)[исходный код] -
Удалить
NodeизGraph. Если вGraphу этого узла ещё есть пользователи, возникает исключение.- Параметры:
-
to_erase (Node) – Узел
Node, который требуется удалить изGraph.
Примечание
Для этого API гарантирована обратная совместимость.
-
find_nodes(*, op, target=None, sort=True)[исходный код] -
Позволяет быстро выполнять поиск узлов
- Параметры:
- Возвращает:
-
Итерируемый объект с узлами, соответствующими заданным операции и цели.
- Тип возвращаемого значения:
Предупреждение
Этот API является экспериментальным и НЕ обеспечивает обратную совместимость.
-
get_attr(qualified_name, type_expr=None)[исходный код] -
Вставить узел
get_attrв граф. Узелget_attrNodeпредставляет получение атрибута из иерархииModule.- Параметры:
-
-
qualified_name (str) – полное квалифицированное имя атрибута, который требуется получить. Например, если трассированный Module содержит подмодуль с именем
foo, у которого есть подмодуль с именемbar, содержащий атрибут с именемbaz, в качествеqualified_nameследует передать полное квалифицированное имяfoo.bar.baz. - type_expr (Optional[Any]) – необязательная аннотация типа, представляющая тип Python, который будет иметь результат этого узла.
-
qualified_name (str) – полное квалифицированное имя атрибута, который требуется получить. Например, если трассированный Module содержит подмодуль с именем
- Возвращает:
-
Недавно созданный и вставленный узел
get_attr. - Тип возвращаемого значения:
Примечание
Для этого метода действуют те же правила точки вставки и выражения типа, что и для
Graph.create_node.Примечание
Для этого API гарантирована обратная совместимость.
-
graph_copy(g, val_map, return_output_node=False)[исходный код] -
Скопировать все узлы из заданного графа в
self.- Параметры:
- Возвращает:
-
Значение в
self, эквивалентное теперь выходному значению вg, если вgбыл узелoutput. В противном случае —None. - Тип возвращаемого значения:
-
tuple[Argument, …] | Sequence[Argument] | Mapping[str, Argument] | slice | range | Node | str | int | float | bool | complex | dtype | Tensor | device | memory_format | layout | OpOverload | SymInt | SymBool | SymFloat | None
Примечание
Для этого API гарантирована обратная совместимость.
-
inserting_after(n=None)[исходный код] -
- Задать точку, в которой 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 гарантирована обратная совместимость.
- Тип возвращаемого значения:
-
_InsertPoint
-
inserting_before(n=None)[исходный код] -
- Задать точку, в которой 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 гарантирована обратная совместимость.
- Тип возвращаемого значения:
-
_InsertPoint
-
lint()[исходный код] -
Выполнить различные проверки этого графа, чтобы убедиться в его корректности. В частности: - Проверить, что узлы принадлежат этому графу - Проверить, что узлы расположены в топологическом порядке - Если у этого графа есть владеющий им GraphModule, проверить, что цели существуют в этом GraphModule
Примечание
Для этого API гарантирована обратная совместимость.
-
materialize_symint(value)[исходный код] -
Удобная обёртка для одного значения вокруг
materialize_symints().Предупреждение
Этот API является экспериментальным и НЕ обеспечивает обратную совместимость.
-
-
materialize_symints(values)[исходный код] -
-
Materialize a list of SymInt/int values as FX subgraphs rooted -
у существующих узлов этого графа, метаданные которых порождают указанные символы (обычно это заполнители SymInt или другие символьные операции).
Предположим, имеется граф с заполнителем тензора
%xформы(s32, s32), и мы хотим записать этот шаг как значение FX, чтобы передать его последующей операции. Есть два способа это сделать.-
g.create_stride_node(%x, 0)создаёт%t = aten.sym_stride.int(%x, 0). Семантика этой операции: «запросить у%xтекущий шаг по измерению 0». Если последующий проход (например, преобразование в формат channels-last для mkldnn) изменит раскладку%x, повторный запускFakeTensorPropперезапишет%t.meta["val"]НОВЫМ шагом. Это правильное поведение, если нужен динамический запрос к производителю. -
g.materialize_symints([%x.meta["val"].stride(0)])обходит символьное выражение SymPys32и создаёт подграф FX, который вычисляет его заново на основе существующего производителяs32(здесь — самого заполнителя%xчерезaten.sym_size.int(%x, 0)или заполнителя SymInt, если он существует). Значение полученного узла — «значениеs32во время выполнения», которое определяется формой входных данных и НЕ ЗАВИСИТ от изменения раскладки%x. Это правильное поведение, если нужно зафиксировать в графе шаг на момент трассировки.
Примечание. Как и другие API создания узлов
Graph(call_function,create_size_nodeи т. д.), узлы добавляются в текущую позицию вставки графа. Позиция вставки по умолчанию (Graph._root.prepend) добавляет новые узлы в конец графа. Если граф уже содержит узелoutput, новые узлы окажутся послеreturn(и останутся несвязанными). Обычно вызывающий код помещает эту операцию вwith graph.inserting_before(graph.output_node()):, чтобы новые узлы попали в тело.Производительность: каждый вызов выполняет сканирование графа для поиска производителей за O(размер графа) и создаёт отдельный для вызова кэш хеш-консолидации
expr_to_proxy. Предпочтительно передавать все необходимые SymInt одним вызовом (или как можно меньшим числом вызовов), а не вызывать функцию отдельно для каждого значения: это позволяет распределить затраты на сканирование графа и объединить SymInt с общими подвыражениями в один подграф. -
Предупреждение
Этот API является экспериментальным и НЕ гарантирует обратную совместимость.
-
-
node_copy(node, arg_transform=<function Graph.<lambda>>)[исходный код] -
Копирует узел из одного графа в другой.
arg_transformнеобходимо преобразовать аргументы из графа узла в граф self. Пример:# Copying all the nodes in `g` into `new_graph` g: torch.fx.Graph = ... new_graph = torch.fx.graph() value_remap = {} for node in g.nodes: value_remap[node] = new_graph.node_copy(node, lambda n: value_remap[n])- Параметры:
- Тип возвращаемого значения:
Примечание
Для этого API гарантирована обратная совместимость.
-
property nodes: _node_list -
Возвращает список узлов, составляющих этот граф.
Обратите внимание, что это
Nodeпредставление списка — двусвязный список. Изменения во время итерации (например, удаление или добавление узла) безопасны.- Возвращает:
-
Двусвязный список узлов. Обратите внимание, что для изменения порядка итерации этого списка можно вызвать
reversed.
-
on_generate_code(make_transformer)[исходный код] -
Регистрирует функцию-преобразователь при генерации кода 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 является экспериментальным и НЕ гарантирует обратную совместимость.
- Тип возвращаемого значения:
-
AbstractContextManager[None]
-
output(result, type_expr=None)[исходный код] -
Вставляет узел
outputNodeвGraph. Узелoutputпредставляет инструкциюreturnв коде Python.result— это возвращаемое значение.- Параметры:
-
- result (Argument) – Возвращаемое значение.
- type_expr (Optional[Any]) – необязательная аннотация типа, представляющая тип Python, который будет иметь выходное значение этого узла.
Примечание
К этому методу применяются те же правила для позиции вставки и выражения типа, что и к
Graph.create_node.Примечание
Для этого API гарантирована обратная совместимость.
-
output_node()[исходный код] -
Предупреждение
Этот API является экспериментальным и НЕ гарантирует обратную совместимость.
- Тип возвращаемого значения:
-
placeholder(name, type_expr=None, default_value)[исходный код] -
Вставляет узел
placeholderв граф.placeholderпредставляет входные данные функции.- Параметры:
-
-
name (str) – Имя входного значения. Оно соответствует имени позиционного аргумента функции, которую представляет этот
Graph. - type_expr (Optional[Any]) – необязательная аннотация типа, представляющая тип Python, который будет иметь выходное значение этого узла. В некоторых случаях она необходима для правильной генерации кода (например, если функция впоследствии используется при компиляции TorchScript).
-
default_value (Any) – Значение по умолчанию для этого аргумента функции. ПРИМЕЧАНИЕ: чтобы использовать
Noneв качестве значения по умолчанию, в этот аргумент следует передатьinspect.Signature.empty, указывая, что параметр _не_ имеет значения по умолчанию.
-
name (str) – Имя входного значения. Оно соответствует имени позиционного аргумента функции, которую представляет этот
- Тип возвращаемого значения:
Примечание
К этому методу применяются те же правила для позиции вставки и выражения типа, что и к
Graph.create_node.Примечание
Для этого API гарантирована обратная совместимость.
-
print_tabular()[исходный код] -
Выводит промежуточное представление графа в табличном формате. Обратите внимание, что для этого API требуется установленный модуль
tabulate.Примечание
Для этого API гарантирована обратная совместимость.
-
process_inputs(*args)[исходный код] -
Обрабатывает аргументы, чтобы их можно было передать графу FX.
Предупреждение
Этот API является экспериментальным и НЕ гарантирует обратную совместимость.
- Тип возвращаемого значения:
-
process_outputs(out)[исходный код] -
Предупреждение
Этот API является экспериментальным и НЕ гарантирует обратную совместимость.
- Тип возвращаемого значения:
-
python_code(root_module, *, verbose=False, include_stride=False, include_device=False, colored=False, expanded_def=False, record_func=False, additional_meta=None)[исходный код] -
Преобразует этот
Graphв допустимый код Python.- Параметры:
-
root_module (str) – Имя корневого модуля, по которому выполняется поиск целевых объектов с полными именами. Обычно это ‘self’.
- Возвращает:
-
src: исходный код Python, представляющий объект; globals: словарь глобальных имён в
src-> объекты, на которые они ссылаются. - Тип возвращаемого значения:
-
Объект PythonCode, состоящий из двух полей
Примечание
Для этого API гарантирована обратная совместимость.
-
set_codegen(codegen)[исходный код] -
Предупреждение
Этот API является экспериментальным и НЕ гарантирует обратную совместимость.
-
-
class torch.fx.Node(graph, name, op, target, args, kwargs, return_type=None)[исходный код] -
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)[исходный код] -
Вставляет
xпосле этого узла в список узлов графа. Эквивалентноself.next.prepend(x)- Параметры:
-
x (Node) – Узел, который нужно поместить после этого узла. Он должен принадлежать тому же графу.
Примечание
Обратная совместимость этого API гарантируется.
-
property args: tuple[tuple[Argument, ...] | Sequence[Argument] | Mapping[str, Argument] | slice | range | Node | str | int | float | bool | complex | dtype | Tensor | device | memory_format | layout | OpOverload | SymInt | SymBool | SymFloat | None, ...] -
Кортеж аргументов этого
Node. Интерпретация аргументов зависит от кода операции узла. Дополнительные сведения см. в строке документацииNode.Это свойство можно присваивать. При присваивании учёт использований и пользователей обновляется автоматически.
-
format_node(placeholder_names=None, maybe_return_typename=None, *, include_tensor_metadata=False)[исходный код] -
Возвращает описательное строковое представление
self.Этот метод можно вызывать без аргументов для отладки.
Эта функция также используется внутри метода
__str__классаGraph. Вместе строки изplaceholder_namesиmaybe_return_typenameобразуют сигнатуру автоматически сгенерированной функцииforwardв окружающем GraphModule этого графа. В противном случаеplaceholder_namesиmaybe_return_typenameиспользовать не следует.- Параметры:
-
-
placeholder_names (list[str] | None) – Список для хранения отформатированных строк, представляющих заполнители в сгенерированной функции
forward. Только для внутреннего использования. -
maybe_return_typename (list[str] | None) – Список из одного элемента для хранения отформатированной строки, представляющей выход сгенерированной функции
forward. Только для внутреннего использования. - include_tensor_metadata (bool) – Нужно ли включать метаданные тензора
-
placeholder_names (list[str] | None) – Список для хранения отформатированных строк, представляющих заполнители в сгенерированной функции
- Возвращает:
-
-
If 1) we’re using format_node as an internal helper -
в методе
__str__классаGraph, и 2) еслиselfявляется узлом-заполнителем, возвращаетNone. В противном случае возвращает описательное строковое представление текущего узла.
-
- Тип возвращаемого значения:
Примечание
Обратная совместимость этого API гарантируется.
-
insert_arg(idx, arg)[исходный код] -
Вставляет позиционный аргумент в список аргументов по заданному индексу.
- Параметры:
-
-
idx (int) – Индекс элемента в
self.args, перед которым нужно выполнить вставку. -
arg (Argument) – Новое значение аргумента, которое нужно вставить в
args
-
idx (int) – Индекс элемента в
Примечание
Обратная совместимость этого API гарантируется.
-
is_impure(impure_random=True)[исходный код] -
Возвращает, является ли эта операция нечистой, то есть является ли её тип placeholder или output либо call_function или call_module, если такая операция нечистая.
- Параметры:
-
impure_random (bool) – Следует ли считать операцию rand нечистой.
- Возвращает:
-
Является ли операция нечистой.
- Тип возвращаемого значения:
Предупреждение
Этот API является экспериментальным и НЕ поддерживает обратную совместимость.
-
property kwargs: dict[str, tuple[Argument, ...] | Sequence[Argument] | Mapping[str, Argument] | slice | range | Node | str | int | float | bool | complex | dtype | Tensor | device | memory_format | layout | OpOverload | SymInt | SymBool | SymFloat | None] -
Словарь именованных аргументов этого
Node. Интерпретация аргументов зависит от кода операции узла. Дополнительные сведения см. в строке документацииNode.Это свойство можно присваивать. При присваивании учёт использований и пользователей обновляется автоматически.
-
property next: Node -
Возвращает следующий
Nodeв связанном списке узлов.- Возвращает:
-
Следующий
Nodeв связанном списке узлов.
-
normalized_arguments(root, arg_types=None, kwarg_types=None, normalize_to_only_use_kwargs=False)[исходный код] -
Возвращает нормализованные аргументы для целей Python. Это означает, что
args/kwargsсопоставляются с сигнатурой модуля/функции, а возвращаются только именованные аргументы в порядке позиций, еслиnormalize_to_only_use_kwargsимеет значение true. Также заполняются значения по умолчанию. Параметры, доступные только по позиции, и параметры varargs не поддерживаются.Поддерживает вызовы модулей.
Для снятия неоднозначности перегрузок могут потребоваться
arg_typesиkwarg_types.- Параметры:
-
- root (torch.nn.Module) – Модуль, относительно которого разрешаются целевые модули.
- arg_types (Optional[Tuple[Any]]) – Кортеж типов позиционных аргументов
- kwarg_types (Optional[Dict[str, Any]]) – Словарь типов именованных аргументов
- normalize_to_only_use_kwargs (bool) – Следует ли нормализовать аргументы так, чтобы использовались только именованные аргументы.
- Возвращает:
-
Возвращает именованный кортеж ArgsKwargsPair или
Noneв случае неудачи. - Тип возвращаемого значения:
-
ArgsKwargsPair | None
Предупреждение
Этот API является экспериментальным и НЕ поддерживает обратную совместимость.
-
prepend(x)[исходный код] -
Вставляет x перед этим узлом в список узлов графа. Пример:
Before: p -> self bx -> x -> ax After: p -> x -> self bx -> ax- Параметры:
-
x (Node) – Узел, который нужно поместить перед этим узлом. Он должен принадлежать тому же графу.
Примечание
Обратная совместимость этого API гарантируется.
-
property prev: Node -
Возвращает предыдущий
Nodeв связанном списке узлов.- Возвращает:
-
Предыдущий
Nodeв связанном списке узлов.
-
replace_all_uses_with(replace_with, delete_user_cb=None, *, propagate_meta=False)[исходный код] -
Заменяет все использования
selfв графе узломreplace_with.- Параметры:
-
-
replace_with (Node) – Узел, которым нужно заменить все использования
self. - delete_user_cb (Callable) – Функция обратного вызова, определяющая, следует ли удалить заданного пользователя узла self.
- propagate_meta (bool) – Следует ли копировать все свойства поля .meta исходного узла в узел-замену. В целях безопасности это допустимо только в том случае, если у узла-замены ещё нет поля .meta.
-
replace_with (Node) – Узел, которым нужно заменить все использования
- Возвращает:
-
Список узлов, в которых было выполнено это изменение.
- Тип возвращаемого значения:
Примечание
Обратная совместимость этого API гарантируется.
-
replace_input_with(old_input, new_input)[исходный код] -
Перебирает входные узлы
selfи заменяет все вхожденияold_inputнаnew_input.- Параметры:
Примечание
Обратная совместимость этого API гарантируется.
-
property stack_trace: str | None -
Возвращает трассировку стека Python, записанную во время трассировки, если она есть. При трассировке с помощью fx.Tracer это свойство обычно заполняется методом
Tracer.create_proxy. Чтобы записывать трассировки стека во время трассировки для отладки, задайтеrecord_stack_traces = Trueв экземпляреTracer. При трассировке с помощью dynamo это свойство по умолчанию заполняется методомOutputGraph.create_proxy.В stack_trace самый внутренний кадр находится в конце строки.
-
update_arg(idx, arg)[исходный код] -
Обновляет существующий позиционный аргумент, задавая ему новое значение
arg. После вызоваself.args[idx] == arg.- Параметры:
-
-
idx (int) – Индекс обновляемого элемента в
self.args -
arg (Argument) – Новое значение аргумента для записи в
args
-
idx (int) – Индекс обновляемого элемента в
Примечание
Обратная совместимость этого API гарантируется.
-
update_kwarg(key, arg)[исходный код] -
Обновляет существующий именованный аргумент, задавая ему новое значение
arg. После вызоваself.kwargs[key] == arg.- Параметры:
-
-
key (str) – Ключ обновляемого элемента в
self.kwargs -
arg (Argument) – Новое значение аргумента для записи в
kwargs
-
key (str) – Ключ обновляемого элемента в
Примечание
Обратная совместимость этого API гарантируется.
-
-
class torch.fx.Tracer(autowrap_modules=(math,), autowrap_functions=())[исходный код] -
Tracer— это класс, реализующий функциональность символьной трассировкиtorch.fx.symbolic_trace. Вызовsymbolic_trace(m)эквивалентенTracer().trace(m).От Tracer можно наследоваться, чтобы переопределять различные аспекты процесса трассировки. Описание аспектов, которые можно переопределить, приведено в строках документации методов этого класса.
Примечание
Для этого API гарантируется обратная совместимость.
-
call_module(m, forward, args, kwargs)[исходный код] -
Метод, определяющий поведение этого
Tracerпри вызове экземпляраnn.Module.По умолчанию проверяется, является ли вызываемый модуль листовым, с помощью
is_leaf_module. Если да, создаётся узелcall_module, ссылающийся наmвGraph. В противном случае обычным образом вызываетсяModule, а операции в его функцииforwardтрассируются.Этот метод можно переопределить, например, чтобы создавать вложенные трассируемые GraphModules или задавать любое другое желаемое поведение при трассировке через границы
Module.- Параметры:
-
- m (Module) – Модуль, для которого создаётся вызов
-
forward (Callable) – Метод forward() объекта
Module, который нужно вызвать - args (Tuple) – Аргументы в точке вызова модуля
- kwargs (Dict) – Именованные аргументы в точке вызова модуля
- Возвращает:
-
Значение, возвращённое при вызове Module. Если был создан узел
call_module, это значение типаProxy. В противном случае возвращается значение, полученное при вызовеModule. - Тип возвращаемого значения:
Примечание
Для этого API гарантируется обратная совместимость.
-
create_arg(a)[исходный код] -
Метод, определяющий поведение трассировки при подготовке значений, используемых в качестве аргументов узлов в
Graph.По умолчанию выполняются следующие действия:
- Перебираются коллекции (например, tuple, list, dict), а для их элементов рекурсивно вызывается
create_args. - Если передан объект Proxy, возвращается ссылка на базовый IR
Node -
Если передан объект Tensor, не являющийся Proxy, для разных случаев создаётся IR:
- Для Parameter создаётся узел
get_attr, ссылающийся на этот Parameter - Tensor, не являющийся Parameter, сохраняется в специальном атрибуте, который ссылается на этот атрибут.
- Для Parameter создаётся узел
Этот метод можно переопределить для поддержки дополнительных типов.
- Параметры:
-
a (Any) – Значение, которое будет создано как
ArgumentвGraph. - Возвращает:
-
Значение
a, преобразованное в соответствующийArgument - Тип возвращаемого значения:
-
Argument
Примечание
Для этого API гарантируется обратная совместимость.
- Перебираются коллекции (например, tuple, list, dict), а для их элементов рекурсивно вызывается
-
create_args_for_root(root_fn, is_module, concrete_args=None)[исходный код] -
Создаёт узлы
placeholder, соответствующие сигнатуре корневого модуляroot. Этот метод анализирует сигнатуру root и создаёт соответствующие узлы, поддерживая также*argsи**kwargs.Предупреждение
Этот API является экспериментальным и НЕ имеет обратной совместимости.
-
create_node(kind, target, args, kwargs, name=None, type_expr=None)[исходный код] -
Вставляет узел графа с заданными target, args, kwargs и name.
Этот метод можно переопределить для дополнительной проверки, валидации или изменения значений, используемых при создании узла. Например, можно запретить запись операций на месте.
Примечание
Для этого API гарантируется обратная совместимость.
- Тип возвращаемого значения:
-
create_proxy(kind, target, args, kwargs, name=None, type_expr=None, proxy_factory_fn=None)[исходный код] -
Создаёт Node из заданных аргументов, а затем возвращает Node, обёрнутый в объект Proxy.
Если kind = ‘placeholder’, создаётся Node, представляющий параметр функции. Если нужно закодировать параметр по умолчанию, используется кортеж
args. В противном случае для узловplaceholderзначениеargsпусто.Примечание
Для этого API гарантируется обратная совместимость.
- Тип возвращаемого значения:
-
get_fresh_qualname(prefix)[исходный код] -
Получает новое имя для префикса и возвращает его. Эта функция гарантирует, что оно не совпадёт с уже существующим атрибутом графа.
Примечание
Для этого API гарантируется обратная совместимость.
- Тип возвращаемого значения:
-
getattr(attr, attr_val, parameter_proxy_cache)[исходный код] -
Метод, определяющий поведение этого
Tracerпри вызове getattr для вызова экземпляраnn.Module.По умолчанию возвращается прокси-значение атрибута. Оно также сохраняется в
parameter_proxy_cache, чтобы при последующих вызовах использовать существующий proxy, а не создавать новый.Этот метод можно переопределить, например, чтобы не возвращать прокси при запросе параметров.
- Параметры:
- Возвращает:
-
Значение, возвращённое вызовом getattr.
- Тип возвращаемого значения:
Предупреждение
Этот API является экспериментальным и НЕ имеет обратной совместимости.
-
is_leaf_module(m, module_qualified_name)[исходный код] -
Метод, определяющий, является ли заданный
nn.Module«листовым» модулем.Листовые модули — это атомарные единицы, которые появляются в IR и на которые ссылаются вызовы
call_module. По умолчанию модули из пространства имён стандартной библиотеки PyTorch (torch.nn) являются листовыми. Все остальные модули трассируются, а их составляющие операции записываются, если для них не задано иное значение этого параметра.- Параметры:
- Тип возвращаемого значения:
Примечание
Для этого API гарантируется обратная совместимость.
-
iter(obj)[исходный код] -
- Вызывается при переборе объекта proxy, например
-
при использовании в потоке управления. Обычно неизвестно, что делать, поскольку значение proxy неизвестно, однако пользовательский трассировщик может добавить к узлу графа дополнительные сведения с помощью create_node и вернуть итератор.
Примечание
Для этого API гарантируется обратная совместимость.
- Тип возвращаемого значения:
-
keys(obj)[исходный код] -
- Вызывается при вызове метода keys() у объекта proxy.
-
Так происходит при вызове ** для proxy. Если ** должен работать в пользовательском трассировщике, этот метод должен возвращать итератор.
Примечание
Для этого API гарантируется обратная совместимость.
- Тип возвращаемого значения:
-
path_of_module(mod)[исходный код] -
Вспомогательный метод для определения полного имени
modв иерархии модулейroot. Например, если уrootесть подмодуль с именемfoo, у которого есть подмодуль с именемbar, то при передачеbarв эту функцию будет возвращена строка «foo.bar».- Параметры:
-
mod (str) –
Module, для которого нужно получить полное имя. - Тип возвращаемого значения:
Примечание
Для этого API гарантируется обратная совместимость.
-
proxy(node)[исходный код] -
Примечание
Для этого API гарантируется обратная совместимость.
- Тип возвращаемого значения:
-
to_bool(obj)[исходный код] -
- Вызывается при преобразовании объекта proxy в логическое значение, например
-
при использовании в потоке управления. Обычно неизвестно, что делать, поскольку значение proxy неизвестно, однако пользовательский трассировщик может добавить к узлу графа дополнительные сведения с помощью create_node и вернуть значение.
Примечание
Для этого API гарантируется обратная совместимость.
- Тип возвращаемого значения:
-
trace(root, concrete_args=None)[исходный код] -
Трассирует
rootи возвращает соответствующее представление FXGraph.rootможет быть экземпляромnn.Moduleили вызываемым объектом Python.Обратите внимание: после этого вызова
self.rootможет отличаться от переданного сюдаroot. Например, если вtrace()передана свободная функция, для использования в качестве корневого объекта и добавления встроенных констант будет создан экземплярnn.Module.- Параметры:
-
-
root (Union[Module, Callable]) – Объект
Moduleили функция, которую нужно трассировать. Для этого параметра гарантируется обратная совместимость. - concrete_args (Optional[Dict[str, any]]) – Конкретные аргументы, которые не следует считать Proxy. Этот параметр является экспериментальным, и его обратная совместимость НЕ гарантируется.
-
root (Union[Module, Callable]) – Объект
- Возвращает:
-
Объект
Graph, представляющий семантику переданногоroot. - Тип возвращаемого значения:
Примечание
Для этого API гарантируется обратная совместимость.
-
-
class torch.fx.Proxy(node, tracer=None)[исходный код] -
Объекты
Proxy— это обёрткиNode, которые передаются по программе во время символьной трассировки и записывают в формируемый граф FX все затрагиваемые ими операции (вызовы функцийtorch, вызовы методов, операторы).При преобразовании графа можно обернуть собственный метод
Proxyвокруг исходногоNode, чтобы использовать перегруженные операторы для добавления дополнительных элементов вGraph.Объекты
Proxyнельзя перебирать. Иными словами, символьный трассировщик выдаст ошибку, еслиProxyиспользуется в цикле или в качестве аргумента функции*args/**kwargs.Есть два основных способа обойти это ограничение: 1. Вынести нетрассируемую логику в функцию верхнего уровня и использовать для неё
fx.wrap. 2. Если поток управления статический (то есть число итераций цикла зависит от некоторого гиперпараметра), код можно оставить на прежнем месте и преобразовать, например, следующим образом:for i in range(self.some_hyperparameter): indexed_item = proxied_value[i]Более подробное описание внутреннего устройства Proxy см. в разделе «Proxy» документации
torch/fx/README.mdПримечание
Для этого API гарантируется обратная совместимость.
-
class torch.fx.Interpreter(module, garbage_collect_values=True, graph=None)[исходный код] -
Интерпретатор выполняет граф 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 is torch.sigmoid: return torch.neg(*args, **kwargs) return super().call_function(target, args, kwargs) 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(target, args, kwargs) def fn(x): return torch.sigmoid(x).neg() gm = torch.fx.symbolic_trace(fn) input = torch.randn(3, 4) result = NegSigmSwapInterpreter(gm).run(input) torch.testing.assert_close(result, torch.neg(input).sigmoid())- Параметры:
-
- module (torch.nn.Module) – Модуль для выполнения
-
garbage_collect_values (bool) – Следует ли удалять значения после их последнего использования при выполнении модуля. Это обеспечивает оптимальное использование памяти во время выполнения. Эту возможность можно отключить, например, чтобы изучить все промежуточные значения выполнения с помощью атрибута
Interpreter.env. -
graph (Optional[Graph]) – Если передан этот аргумент, интерпретатор выполнит указанный граф вместо
module.graph, используя предоставленный аргументmoduleдля обработки запросов состояния.
Примечание
Обратная совместимость этого API гарантируется.
-
boxed_run(args_list)[исходный код] -
Выполняет
moduleс помощью интерпретации и возвращает результат. Используется соглашение о вызове «boxed»: аргументы передаются в виде списка, который интерпретатор очистит. Это обеспечивает своевременное освобождение входных тензоров.Примечание
Обратная совместимость этого API гарантируется.
- Тип возвращаемого значения:
-
call_function(target, args, kwargs)[исходный код] -
Выполняет узел
call_functionи возвращает результат.- Параметры:
-
- target (Target) – Цель вызова для этого узла. Подробности семантики см. в описании Node
- args (Tuple) – Кортеж позиционных аргументов этого вызова
- kwargs (Dict) – Словарь именованных аргументов этого вызова
- Тип возвращаемого значения:
- Возвращает
-
Any: значение, возвращённое при вызове функции
Примечание
Обратная совместимость этого API гарантируется.
-
call_method(target, args, kwargs)[исходный код] -
Выполняет узел
call_methodи возвращает результат.- Параметры:
-
- target (Target) – Цель вызова для этого узла. Подробности семантики см. в описании Node
- args (Tuple) – Кортеж позиционных аргументов этого вызова
- kwargs (Dict) – Словарь именованных аргументов этого вызова
- Тип возвращаемого значения:
- Возвращает
-
Any: значение, возвращённое при вызове метода
Примечание
Обратная совместимость этого API гарантируется.
-
call_module(target, args, kwargs)[исходный код] -
Выполняет узел
call_moduleи возвращает результат.- Параметры:
-
- target (Target) – Цель вызова для этого узла. Подробности семантики см. в описании Node
- args (Tuple) – Кортеж позиционных аргументов этого вызова
- kwargs (Dict) – Словарь именованных аргументов этого вызова
- Тип возвращаемого значения:
- Возвращает
-
Any: значение, возвращённое при вызове модуля
Примечание
Обратная совместимость этого API гарантируется.
-
fetch_args_kwargs_from_env(n)[исходный код] -
Получает конкретные значения
argsиkwargsузлаnиз текущего окружения выполнения.- Параметры:
-
n (Node) – Узел, для которого следует получить
argsиkwargs. - Возвращает:
-
argsиkwargsс конкретными значениями дляn. - Тип возвращаемого значения:
-
Tuple[Tuple, Dict]
Примечание
Обратная совместимость этого API гарантируется.
-
fetch_attr(target)[исходный код] -
Получает атрибут из иерархии
Moduleобъектаself.module.- Параметры:
-
target (str) – Полное имя атрибута, который необходимо получить
- Возвращает:
-
Значение атрибута.
- Тип возвращаемого значения:
-
Any
Примечание
Обратная совместимость этого API гарантируется.
-
get_attr(target, args, kwargs)[исходный код] -
Выполняет узел
get_attr. Извлекает значение атрибута из иерархииModuleобъектаself.module.- Параметры:
-
- target (Target) – Цель вызова для этого узла. Подробности семантики см. в описании Node
- args (Tuple) – Кортеж позиционных аргументов этого вызова
- kwargs (Dict) – Словарь именованных аргументов этого вызова
- Возвращает:
-
Полученное значение атрибута
- Тип возвращаемого значения:
-
Any
Примечание
Обратная совместимость этого API гарантируется.
-
map_nodes_to_values(args, n)[исходный код] -
Рекурсивно проходит по
argsи находит конкретное значение для каждогоNodeв текущем окружении выполнения.- Параметры:
-
- args (Argument) – Структура данных, в которой нужно найти конкретные значения
-
n (Node) – Узел, которому принадлежит
args. Используется только для формирования сообщений об ошибках.
- Тип возвращаемого значения:
-
tuple[Argument, …] | Sequence[Argument] | Mapping[str, Argument] | slice | range | Node | str | int | float | bool | complex | dtype | Tensor | device | memory_format | layout | OpOverload | SymInt | SymBool | SymFloat | None
Примечание
Обратная совместимость этого API гарантируется.
-
output(target, args, kwargs)[исходный код] -
Выполняет узел
output. По сути, просто получает значение, на которое ссылается узелoutput, и возвращает его.- Параметры:
-
- target (Target) – Цель вызова для этого узла. Подробности семантики см. в описании Node
- args (Tuple) – Кортеж позиционных аргументов этого вызова
- kwargs (Dict) – Словарь именованных аргументов этого вызова
- Возвращает:
-
Возвращаемое значение, на которое ссылается выходной узел
- Тип возвращаемого значения:
-
Any
Примечание
Обратная совместимость этого API гарантируется.
-
placeholder(target, args, kwargs)[исходный код] -
Выполняет узел
placeholder. Обратите внимание, что этот метод изменяет состояние:Interpreterподдерживает внутренний итератор по аргументам, переданным вrun, и этот метод возвращает next() для этого итератора.- Параметры:
-
- target (Target) – Цель вызова для этого узла. Подробности семантики см. в описании Node
- args (Tuple) – Кортеж позиционных аргументов этого вызова
- kwargs (Dict) – Словарь именованных аргументов этого вызова
- Возвращает:
-
Полученное значение аргумента.
- Тип возвращаемого значения:
-
Any
Примечание
Обратная совместимость этого API гарантируется.
-
run(*args, initial_env=None, enable_io_processing=True)[исходный код] -
Выполняет
moduleс помощью интерпретации и возвращает результат.- Параметры:
-
- *args (Any) – Аргументы модуля для выполнения, в позиционном порядке
-
initial_env (Optional[Dict[Node, Any]]) – Необязательное начальное окружение выполнения. Это словарь, сопоставляющий
Nodeс любым значением. Например, его можно использовать для предварительного заполнения результатов для определённыхNodes, чтобы выполнить в интерпретаторе только частичное вычисление. - enable_io_processing (bool) – Если значение истинно, перед использованием сначала обрабатываются входные и выходные данные с помощью функций process_inputs и process_outputs графа.
- Возвращает:
-
Значение, возвращённое при выполнении модуля
- Тип возвращаемого значения:
-
Any
Примечание
Обратная совместимость этого API гарантируется.
-
run_node(n)[исходный код] -
Выполняет определённый узел
nи возвращает результат. В зависимости отnode.opвызывает placeholder, get_attr, call_function, call_method, call_module или output.- Параметры:
-
n (Node) – Узел для выполнения
- Возвращает:
-
Результат выполнения
n - Тип возвращаемого значения:
-
Any
Примечание
Обратная совместимость этого API гарантируется.
-
class torch.fx.Transformer(module)[исходный код] -
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 is torch.sigmoid: return torch.neg(*args, **kwargs) return super().call_function(target, args, kwargs) 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(target, args, kwargs) def fn(x): return torch.sigmoid(x).neg() gm = torch.fx.symbolic_trace(fn) transformed: torch.nn.Module = NegSigmSwapXformer(gm).transform() input = torch.randn(3, 4) torch.testing.assert_close(transformed(input), torch.neg(input).sigmoid())- Параметры:
-
module (GraphModule) –
Moduleдля преобразования.
Примечание
Обратная совместимость этого API гарантируется.
-
call_function(target, args, kwargs)[исходный код] -
Примечание
Обратная совместимость этого API гарантируется.
- Тип возвращаемого значения:
-
call_module(target, args, kwargs)[исходный код] -
Примечание
Обратная совместимость этого API гарантируется.
- Тип возвращаемого значения:
-
get_attr(target, args, kwargs)[исходный код] -
Выполняет узел
get_attr. ВTransformerэтот метод переопределён для вставки нового узлаget_attrв выходной граф.- Параметры:
-
- target (Target) – Цель вызова для этого узла. Подробности семантики см. в описании Node
- args (Tuple) – Кортеж позиционных аргументов этого вызова
- kwargs (Dict) – Словарь именованных аргументов этого вызова
- Тип возвращаемого значения:
Примечание
Обратная совместимость этого API гарантируется.
-
placeholder(target, args, kwargs)[исходный код] -
Выполняет узел
placeholder. ВTransformerэтот метод переопределён для вставки нового узлаplaceholderв выходной граф.- Параметры:
-
- target (Target) – Цель вызова для этого узла. Подробности семантики см. в описании Node
- args (Tuple) – Кортеж позиционных аргументов этого вызова
- kwargs (Dict) – Словарь именованных аргументов этого вызова
- Тип возвращаемого значения:
Примечание
Обратная совместимость этого API гарантируется.
-
transform()[исходный код] -
Преобразует
self.moduleи возвращает преобразованныйGraphModule.Примечание
Обратная совместимость этого API гарантируется.
- Тип возвращаемого значения:
-
torch.fx.replace_pattern(gm, pattern, replacement)[исходный код] -
Находит все возможные непересекающиеся наборы операторов и их зависимостей по данным (
pattern) в графе GraphModule (gm), а затем заменяет каждый найденный подграф другим подграфом (replacement).- Параметры:
-
- gm (GraphModule) – GraphModule, содержащий граф, с которым нужно работать
-
pattern (Callable[[...], Any] | GraphModule) – Подграф, который нужно найти в
gmдля замены -
replacement (Callable[[...], Any] | GraphModule) – Подграф, которым нужно заменить
pattern
- Возвращает:
-
Список объектов
Match, представляющих места в исходном графе, которым соответствуетpattern. Если совпадений нет, список пуст. ОпределениеMatch:class Match(NamedTuple): # Node from which the match was found anchor: Node # Maps nodes in the pattern subgraph to nodes in the larger graph nodes_map: Dict[Node, Node] - Тип возвращаемого значения:
-
List[Match]
Примеры:
import torch from torch.fx import symbolic_trace, subgraph_rewriter class M(torch.nn.Module): def __init__(self) -> None: 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]) 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. Сопоставление с шаблоном выполняется на основе связей использования и определения, а не имён узлов. Например, если вpatternестьp = torch.cat([a, b]), можно найтиm = torch.cat([a, b])в исходной функцииforward, несмотря на различие имён переменных (pиm).Оператор
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 гарантируется обратная совместимость.
-
torch.fx.traceback.annotate(annotation_dict)[исходный код] -
Временно добавляет пользовательские аннотации в текущий контекст трассировки. Узел fx_node, созданный в этом контексте трассировки, будет содержать пользовательские аннотации в поле node.metadata[“custom”].
Этот менеджер контекста позволяет добавлять произвольные метаданные в систему трассировки PT2, обновляя глобальный словарь
current_meta[“custom”]. Аннотации автоматически отменяются при выходе из контекста.Узлы накопления градиента не будут аннотированы.
Этот API предназначен для опытных пользователей, которым необходимо добавлять дополнительные метаданные к узлам fx (например, для отладки, анализа или внешних инструментов) во время трассировки при экспорте.
Примечание
Обратная совместимость этого API не гарантируется; он может измениться в будущих выпусках.
Примечание
Этот API несовместим с fx.symbolic_trace и jit.trace. Он предназначен для использования с семейством трассировщиков PT2, например torch.export и dynamo.
- Параметры:
-
annotation_dict (dict) – Словарь пользовательских пар «ключ-значение» для добавления в метаданные трассировки FX.
- Тип возвращаемого значения:
-
Iterator[None]
Пример
После выхода из контекста пользовательские аннотации удаляются.
>>> with annotate({"source": "custom_pass", "tag": 42}): ... pass # Your computation hereПредупреждение
Этот API является экспериментальным, и его обратная совместимость НЕ гарантируется.
-
torch.fx.passes.tools_common.stable_topological_sort(gm)[исходный код] -
Заменяет граф заданного GraphModule графом, содержащим те же узлы, что и исходный, но расположенные в топологическом порядке с максимально возможным сохранением исходного порядка узлов.
Эта функция выполняет устойчивую топологическую сортировку, при которой узлы располагаются в порядке, который: 1. Соблюдает зависимости по данным (топологический порядок). 2. Сохраняет исходный порядок узлов, если ограничения зависимостей отсутствуют.
Алгоритм использует алгоритм Кана с очередью с приоритетами: узлы, у которых выполнены все зависимости, добавляются в минимальную кучу и упорядочиваются по исходной позиции. Это гарантирует, что среди готовых узлов мы всегда обрабатываем узел, который раньше всех располагался в исходном порядке.
- Параметры:
-
gm (GraphModule) – Модуль графа для топологической сортировки. Он изменяется на месте.
- Возвращает:
-
Модуль графа, отсортированный на месте
- Тип возвращаемого значения:
Предупреждение
Этот API является экспериментальным, и его обратная совместимость НЕ гарантируется.
torch.fx.annotate
annotate
| Аннотирует объект Proxy заданным типом. |
torch.fx.node
has_side_effect
| Регистрирует функцию, которую нельзя удалять как мёртвый код с помощью fx.graph.eliminate_dead_code |
map_aggregate
| Рекурсивно применяет fn к каждому объекту, содержащемуся в arg. |
map_arg
| Рекурсивно применяет fn к каждому узлу Node, содержащемуся в arg. |
torch.fx.operator_schemas
check_for_mutable_operation
| |
create_type_hint
| Создаёт подсказку типа для заданного аргумента. |
get_signature_for_torch_op
| Для оператора из пространства имён |
normalize_function
| Возвращает нормализованные аргументы функций PyTorch. |
normalize_module
| Возвращает нормализованные аргументы модулей PyTorch. |
type_matches
|
torch.fx.traceback
annotate_fn
| Декоратор, оборачивающий функцию в менеджер контекста annotate. |
get_graph_provenance_json
| Для fx.Graph возвращает json с информацией о происхождении каждого узла. |
NodeSource
| NodeSource — это структура данных, содержащая информацию о происхождении узла. |
NodeSourceAction
| Перечисление, представляющее действие, выполненное для создания узла при отслеживании происхождения. |
torch.fx.subgraph_rewriter
replace_pattern
| Находит все возможные непересекающиеся наборы операторов и их зависимостей по данным ( |
replace_pattern_with_filters
| Описание см. в документации replace_pattern. |
torch.fx.tensor_type
is_consistent
| Бинарное отношение, обозначаемое символом ~, определяющее, согласуется ли t1 с t2. |
is_more_precise
| Бинарное отношение, обозначаемое символом <=, определяющее, является ли t1 более точным, чем t2. |
torch.fx.passes.backends.cudagraphs
partition_cudagraphs
| Разбивает граф FX на подмодули GraphModule, которые могут корректно выполняться с использованием CUDA Graphs. |
torch.fx.passes.graph_manipulation
-
torch.fx.passes.graph_manipulation.get_size_of_all_nodes(fx_module, args=None)[исходный код] -
- Для модуля графа fx обновляет каждый узел, указывая его общий размер (веса + смещение + выход)
-
и размер выходных данных (output_size). Для узла, не являющегося модулем, общий размер равен размеру выходных данных. Возвращает общий размер.
Предупреждение
Этот API является экспериментальным, и его обратная совместимость НЕ гарантируется.
-
torch.fx.passes.graph_manipulation.get_size_of_node(fx_module, node)[исходный код] -
- Для узла с node.dtype и node.shape возвращает его общий размер и размер выходных данных.
-
total_size = weights + bias + output_size
Предупреждение
Этот API является экспериментальным, и его обратная совместимость НЕ гарантируется.
- Тип возвращаемого значения:
-
size_bytes
replace_target_nodes_with
| Изменяет все узлы в fx_module.graph.nodes, соответствующие указанным коду операции и цели, и обновляет их, чтобы они соответствовали новому коду операции и цели. |
torch.fx.passes.infra.pass_manager
pass_result_wrapper
| Обёртка для проходов, которые в данный момент не возвращают PassResult. |
this_before_that_pass_constraint
| Задаёт частичный порядок (функцию «зависит от»), в котором |
torch.fx.passes.operator_support
any_chain
| Объединяет последовательность экземпляров |
chain
| Объединяет последовательность экземпляров |
create_op_support
| Оборачивает функцию |
torch.fx.passes.param_fetch
-
torch.fx.passes.param_fetch.default_matching(name, target_version)[исходный код] -
Метод сопоставления по умолчанию
Предупреждение
Этот API является экспериментальным, и его обратная совместимость НЕ гарантируется.
- Тип возвращаемого значения:
-
torch.fx.passes.param_fetch.extract_attrs_for_lowering(mod)[исходный код] -
-
If mod is in module_fetch_book, fetch the mod’s attributes that in the module_fetch_book -
после проверки совместимости версии модуля с
module_fetch_book.
Предупреждение
Этот API является экспериментальным, и его обратная совместимость НЕ гарантируется.
-
-
torch.fx.passes.param_fetch.lift_lowering_attrs_to_nodes(fx_module)[исходный код] -
Рекурсивно обходит все узлы
fx_moduleи получает атрибуты модуля, если узел является листовым модулем.Предупреждение
Этот API является экспериментальным, и его обратная совместимость НЕ гарантируется.
torch.fx.passes.pass_manager
inplace_wrapper
| Вспомогательная обёртка для проходов, изменяющих объект на месте. |
log_hook
| Записывает в журнал результат вызываемого объекта. |
loop_pass
| Вспомогательная обёртка для проходов, которые необходимо применять несколько раз. |
these_before_those_pass_constraint
| Задаёт частичный порядок (функцию «зависит от»), в котором |
this_before_that_pass_constraint
| Задаёт частичный порядок (функцию «зависит от»), в котором |
torch.fx.passes.regional_inductor
regional_inductor
| Выделяет области, помеченные для inductor, и компилирует их с помощью inductor. |
torch.fx.passes.reinplace
reinplace
| Для заданного fx.GraphModule изменяет его, выполняя «reinplacing» — мутацию узлов графа. |
torch.fx.passes.split_utils
-
torch.fx.passes.split_utils.split_by_tags(gm, tags, return_fqn_mapping=False, return_tuple=False, GraphModuleCls=<class 'torch.fx.graph_module.GraphModule'>)[исходный код] -
Разбивает GraphModule, используя теги узлов графа. Порядок тегов сохраняется. Например, если заданы теги = [“a”, “b”, “c”], функция создаст начальные подмодули в порядке “a”, “b”, “c”.
Чтобы задать тег:
gm.graph.nodes[idx].tag = "mytag"
В результате все узлы с одинаковым тегом будут извлечены и помещены в собственный подмодуль. Для узлов placeholder, output и get_attr тег игнорируется. Узлы placeholder и output создаются при необходимости, а узлы get_attr копируются в подмодули, где они используются.
Для следующего определения модуля:
class SimpleModule(torch.nn.Module): def __init__(self) -> None: super().__init__() self.linear1 = torch.nn.Linear(...) self.linear2 = torch.nn.Linear(...) self.linear3 = torch.nn.Linear(...) def forward(self, in1, in2): r1 = self.linear1(in1) r2 = self.linear2(in2) r3 = torch.cat([r1, r2]) return self.linear3(r3)Если пометить узел, соответствующий in1, тегом sc.REQUEST_ONLY.lower(), получится следующее разбиение:
ro:
def forward(self, in1): self = self.root linear1 = self.linear1(in1) return linear1main:
def forward(self, in2, linear1): self = self.root linear2 = self.linear2(in2) cat_1 = torch.cat([linear1, linear2]) linear3 = self.linear3(cat_1) return linear3main:
def forward(self, in1, in2): self = self.root ro_0 = self.ro_0(in1) main_1 = self.main_1(in2, ro_0) return main_1- Возвращает:
-
Граф torch fx после разбиения
- orig_to_split_fqn_mapping: отображение исходного fqn в fqn
-
после разбиения для call_module и get_attr.
- Тип возвращаемого значения:
-
split_gm
Предупреждение
Этот API является экспериментальным, и его обратная совместимость НЕ гарантируется.
torch.fx.passes.tools_common
is_node_output_tensor
| Проверяет, возвращает ли узел на выходе Tensor. |
torch.fx.passes.utils.common
-
torch.fx.passes.utils.common.lift_subgraph_as_module(gm, subgraph, comp_name='', class_name='GraphModule')[исходный код] -
Создать GraphModule для подграфа, скопировав необходимые атрибуты из исходного родительского graph_module.
- Параметры:
-
- gm (GraphModule) – родительский модуль графа
-
subgraph (
torch.fx.Graph) – допустимый подграф, содержащий скопированные узлы из родительского графа - comp_name (str) – имя нового компонента
- class_name (str) – имя подмодуля
- Тип возвращаемого значения:
-
tuple[GraphModule, dict[str, str]]
Предупреждение
Этот API является экспериментальным и НЕ обладает обратной совместимостью.
compare_graphs
| Возвращает True, если два графа идентичны, то есть |
torch.fx.passes.utils.fuser_utils
-
torch.fx.passes.utils.fuser_utils.fuse_as_graphmodule(gm, nodes, module_name, partition_lookup_table=None, *, always_return_tuple=False)[исходный код] -
Объединить узлы в graph_module в GraphModule.
- Параметры:
-
- gm (GraphModule) – целевой graph_module
-
nodes (List[Node]) – список узлов в
gmдля объединения; узлы должны быть отсортированы топологически - module_name (str) – имя класса объединённого GraphModule
- partition_lookup_table (Optional[Dict[Node, None]]) – необязательный словарь узлов для ускорения поиска
- always_return_tuple (bool) – всегда ли возвращать кортеж, даже если выход только один
- Возвращает:
-
объединённый модуль графа, узел которого является копией
nodesвgmoriginal_inputs (Tuple[Node, …]): входные узлы для
nodesв исходномgmoriginal_outputs (Tuple[Node, …]): узлы-потребители
nodesв исходномgm - Тип возвращаемого значения:
-
fused_gm (GraphModule)
Предупреждение
Этот API является экспериментальным и НЕ обладает обратной совместимостью.
torch.fx.passes.utils.source_matcher_utils
-
torch.fx.passes.utils.source_matcher_utils.get_source_partitions(graph, wanted_sources, filter_fn=None)[исходный код] -
- Параметры:
- Возвращает:
-
Словарь, сопоставляющий переданные источники со списком SourcePartition, соответствующих списку узлов, декомпозированных из указанного источника.
- Тип возвращаемого значения:
Предупреждение
Этот API является экспериментальным и НЕ обладает обратной совместимостью.
check_subgraphs_connected
| Для двух подграфов A и B (представленных в виде списков узлов) проверяет, есть ли в A узлы, соединённые хотя бы с одним узлом в B, то есть существует ли узел в B, который использует узел из A (но не наоборот). |
get_unique_attr_name_in_module
| Проверяет, что имя уникально (в модуле) и может представлять атрибут. |
split_const_subgraphs
| Просматривает |
-
torch.fx.passes.annotate_getitem_nodes.annotate_getitem_nodes(graph)[исходный код] -
Аннотирует тип узлов getitem, определяемый по типу узла последовательности. Если узел последовательности не аннотирован типом, ничего не делает. В настоящее время поддерживает узлы getitem для узлов последовательностей tuple, list и NamedTuple.
Это полезно, поскольку аннотации локальных имён внутри функции теряются при преобразованиях FX. Добавление известных аннотаций типов обратно к узлам getitem повышает совместимость со скриптами JIT.
- Параметры:
-
graph (
torch.fx.Graph) – Граф, который нужно аннотировать
insert_deferred_runtime_asserts
| Во время трассировки можно обнаружить, что для некоторых значений, зависящих от данных, предусмотрена проверка во время выполнения; например, torch.empty(x.item()) подразумевает проверку во время выполнения, что x.item() >= 0. |
TensorMetadata
| Структура, содержащая важную информацию о тензоре в программе PyTorch. |
-
torch.fx.passes.split_utils.move_non_tensor_nodes_on_boundary(subgraphs)[исходный код] -
Перемещает узлы, не являющиеся тензорами, на границу между подграфами.
Для каждого подграфа:
- Находит узлы, тип которых не является тензором и у которых есть потомки в другом подграфе, и помещает их в очередь для следующего шага
-
Выполняет BFS для узлов в очереди и DFS для каждого узла; допустим, это узел X, находящийся в подграфе A:
- если он находится в to_subgraph, возвращает результат (продолжает DFS)
- если он находится в from_subgraph, добавляет узлы в nodes_to_move и продолжает DFS
- в противном случае это означает, что его нельзя переместить
- также проверяет, нужно ли добавить родительский узел X в очередь. (В очереди могут быть повторяющиеся узлы; каждый узел обрабатывается только один раз)
- Параметры:
-
subgraphs (list[Subgraph]) – Список подграфов, содержащих узлы для обработки
Предупреждение
Этот API является экспериментальным и НЕ обладает обратной совместимостью.
-
torch.fx.passes.splitter_base.generate_inputs_for_submodules(model, inputs, target_submodules, deepcopy=False)[исходный код] -
Формирует входные данные для целевых подмодулей заданной модели. Обратите внимание: эта функция не работает, если два подмодуля ссылаются на один и тот же объект.
- Параметры:
- Возвращает:
-
Словарь, сопоставляющий имена подмодулей с их входными данными.
- Тип возвращаемого значения:
Предупреждение
Этот API является экспериментальным и НЕ обладает обратной совместимостью.
NodeEvent
| Событие, произошедшее с узлом при разделении графа. |
NodeEventTracker
| Отслеживает события узлов во время выполнения разделителя. |
SubgraphMatcherWithNameNodeMap
| Расширяет SubgraphMatcher, добавляя поддержку поиска узлов сопоставленного подграфа по имени узла, |
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/fx.html