Написание преобразований графа для ATen IR
Создано: 11 июня 2025 г. | Последнее обновление: 3 декабря 2025 г.
Проходы
Поскольку ATen IR находится на уровне FX Graph/GraphModule, любые преобразования, написанные для графов FX, можно легко применять к ATen IR. Если вы знакомы с написанием преобразований графов FX, здесь всё будет так же.
Самый простой способ написать преобразование — пройтись по заданному графу и напрямую изменить узлы внутри него.
Например, допустим, мы хотим заменить вызовы torch.ops.aten.add.Tensor() вызовами torch.ops.aten.mul.Tensor():
import torch
def replace_add_with_mul(gm: torch.fx.GraphModule) -> torch.fx.GraphModule:
for node in gm.graph.nodes:
if node.op == "call_function" and node.target == torch.ops.aten.add.Tensor:
node.target = torch.ops.aten.mul.Tensor
Мы также можем удалять и добавлять новые узлы с помощью вспомогательных функций FX, описанных в документации Graph. Например, если мы хотим вставить torch.ops.aten.relu.default() после вызова add:
import torch
def insert_relu_after_add(gm: torch.fx.GraphModule) -> torch.fx.GraphModule:
for node in gm.graph.nodes:
if node.op == "call_function" and node.target == torch.ops.aten.add.Tensor:
# Specifies the insertion point. Any nodes added to the graph within
# this scope will be inserted after `node`
with gm.graph.inserting_after(node):
# Insert a new `call_function` node with op `torch.ops.aten.relu.default`
new_relu_node = gm.graph.call_function(torch.ops.aten.relu.default, args=(node,))
# Replace all the places that use `node` to now use the `new_relu_node`
node.replace_all_uses_with(new_relu_node)
В целом преобразования можно условно разделить по нескольким осям:
Ось A: 1. Создание отображения «один ко многим» (например, декомпозиция) 2. Создание отображения «многие к одному» (например, слияние)
Ось B: 1. Прямой обход (например, распространение форм) 2. Обратный обход (например, удаление мёртвого кода)
Ось C: 1. Зависимость от локальной информации об узле (например, преобразование в вариант с выходом) 2. Зависимость от глобальной информации о графе (например, планирование памяти)
Мы предполагаем, что эти сценарии использования встречаются с такой частотой: 1. A.1, B.1, C.1 2. A.2 3. B.2, C.2
Хотя все преобразования графа можно выполнять путём непосредственного изменения графа, для удобства мы также предоставляем вспомогательные инструменты для сценариев использования уровней 1 и 2.
Transformer
Для сценариев использования уровня 1 (создание отображений «один ко многим», прямой обход и работа с локальной информацией об узлах) можно использовать класс Transformer, чтобы выполнить каждый узел и заново создать граф, применив указанные преобразования.
Преобразование «один к одному»
Например, чтобы заменить операцию A другой операцией B, можно запустить GraphModule и каждый раз, когда встречается операция A, возвращать операцию B.
Пример:
class ReplaceAddWithMul(torch.fx.Transformer):
def call_function(self, target, args, kwargs):
if target != torch.ops.aten.add.Tensor:
return super().call_function(target, args, kwargs)
return super().call_function(torch.ops.aten.mul.Tensor, args, kwargs)
transformed_graph_module = ReplaceAddWithMul(graph_module).transform()
Вызов super().call_function(target, args, kwargs, meta) создаёт узел FX типа call_function и возвращает результат выполнения оператора с заданными аргументами.
Преобразование «один ко многим»
Для отображения «один ко многим», например замены операции A двумя другими операциями B и C, нужно дважды вызвать super().call_function, чтобы создать два узла FX: один с операцией B, другой с операцией C. Затем следует вернуть результат выполнения операции C.
Например:
class ReplaceAddWithMulSub(torch.fx.Transformer):
"""
Original:
def f(x, y):
return x + y
After pass:
def f(x, y):
z = x * y
return z - y
"""
def call_function(self, target, args, kwargs):
if target != torch.ops.aten.add.Tensor:
return super().call_function(target, args, kwargs)
x, y = args
mul_res = super().call_function(torch.ops.aten.mul.Tensor, args, {})
return super().call_function(torch.ops.aten.sub.Tensor, (mul_res, y), {})
transformed_graph_module = ReplaceAddWithMulSub(graph_module).transform()
Преобразование «один в ноль»
Чтобы удалить операцию, достаточно вернуть значение, переданное функции:
class RemoveDetachPass(torch.fx.Transformer):
def call_function(self, target, args, kwargs):
if target not in (
torch.ops.aten.detach.default,
torch.ops.aten.detach_copy.default,
):
return super().call_function(target, args, kwargs, meta)
assert len(args) == 1
return args[0]
transformed_graph_module = RemoveDetachPass(graph_module).transform()
Использование локальной информации
В качестве примера использования локальной информации об узлах можно преобразовать все скалярные значения в графе в тензоры: запустить заданный fx.GraphModule и преобразовать в тензор каждый аргумент, содержащий скаляр. Это может выглядеть так:
def args_map(target, fn, args, kwargs):
assert isinstance(args, tuple)
assert isinstance(kwargs, dict)
args = list(args)
kwargs = kwargs.copy()
# Update the argument based on the function passed
def update(key, args, schema):
args[key] = fn(args[key], schema)
# Update each argument in the schema
for i, schema in enumerate(target._schema.arguments):
if schema.name in kwargs:
update(schema.name, kwargs, schema)
elif not schema.kwarg_only and i < len(args):
update(i, args, schema)
return tuple(args), kwargs
class ScalarToTensorPass(torch.fx.Transformer):
def call_function(self, target, args, kwargs):
breakpoint()
def try_coerce(value, arg):
return (
torch.tensor(value)
if isinstance(value, (float, int, bool))
and type(arg.type) == torch.TensorType
else value
)
args, kwargs = args_map(target, try_coerce, args, kwargs)
return super().call_function(target, args, kwargs)
transformed_graph_module = ScalarToTensorPass(graph_module).transform()
Переписыватель подграфов
Для создания отображений «многие к одному» можно использовать переписыватель подграфов FX. Получив pattern, он создаёт подграф операторов, соответствующих шаблону, а затем заменяет каждый найденный подграф на replacement.
Примечание:
This is an inplace operation.
Входные данные pattern и replacement должны быть вызываемыми функциями или GraphModule, содержащими те же операторы, что используются в графе (операции ATen), чтобы переписыватель подграфов мог найти в графе нужный шаблон. Входные данные вызываемых объектов шаблона и замены при сопоставлении рассматриваются как шаблонные параметры.
Пример:
from torch.fx import subgraph_rewriter
def replace_patterns(graph_module):
def pattern(x, y):
x = torch.ops.aten.add.Tensor(x, y)
x = torch.ops.aten.mul.Tensor(x, y)
return x
def replacement(x, y):
return torch.ops.aten.sub.Tensor(x, y)
replaced_patterns = subgraph_rewriter.replace_pattern_with_filters(
traced_module, pattern, replacement
)
Переписыватель подграфов возвращает список ReplacedPatterns:
@dataclass
class ReplacedPatterns:
# 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 of nodes that were added into the graph
replacements: List[Node]
Примечание:
The nodes created by the subgraph rewriter will not have the metadata that is populated in the matched nodes, but you can use `ReplacedPatterns.nodes_map` to find the nodes in the original graph that were matched, and `ReplacedPatterns.replacements` to find the nodes that were replaced in the transformed graph.
Менеджер проходов
PassManager — это класс для запуска нескольких проходов над заданным модулем графа. При создании экземпляра PassManager ему передаётся список проходов, которые нужно запустить, а также задаются несколько флагов. Чтобы запустить набор проходов для модуля графа, достаточно передать модуль графа непосредственно экземпляру PassManager.
Пример:
from torch.fx.passes.infra.pass_manager import PassManager
pm = PassManager(
passes=[replace_add_with_div, replace_div_with_mul],
run_checks_after_each_pass=True,
suppress_check_failures=False,
)
graph_module_out = pm(graph_module)
Чтобы добавить общий набор проверок, выполняемых после каждого прохода, можно вызвать функцию set_checks(check: Callable), принимающую в качестве входных данных вызываемую функцию. Если установлен флаг run_checks_after_each_pass, функция check будет вызываться после выполнения каждого прохода над модулем графа.
Пример:
pm = PassManager(passes=[replace_add_with_div, replace_div_with_mul])
def check_div_target(graph_module):
for node in graph_module.graph.nodes:
if node.op == "call_function" and node.target != torch.div:
raise ValueError("Target should be div!")
pm.add_checks(check_div_target)
pm(graph_module) # raises ValueError after replace_div_with_mul pass
Разбиение на разделы
Для разбиения графа можно использовать несколько распространённых разделителей на основе графов FX.
Сопоставитель подграфов
Чтобы найти в графе подграфы, соответствующие определённому шаблону, можно использовать SubgraphMatcher FX.
Атрибуты класса:
-
pattern (Graph): шаблон для сопоставления. Узлы-заполнители в графе при сопоставлении рассматриваются как шаблонные параметры. -
match_output (bool): если True, выходной узел в графе шаблона считается частью целевого шаблона. Если False, выходной узел игнорируется при сопоставлении. -
match_placeholder (bool): если True, узел-заполнитель в графе шаблона считается частью целевого шаблона. Если False, узлы-заполнители используются как шаблонные параметры. -
remove_overlapping_matches (bool): если True, при наличии пересекающихся совпадений будет возвращено только первое. -
ignore_literals (bool): если True, проверка равенства литералов выполняться не будет, и вместо этого они будут считаться шаблонными параметрами.
Пример:
from torch.fx.passes.utils.matcher_utils import SubgraphMatcher
class LargeModel(torch.nn.Module):
def __init__(self):
super().__init__()
self._weight = torch.nn.Parameter(torch.ones(3, 3))
self._bias = torch.nn.Parameter(torch.ones(3, 3))
def forward(self, x):
return torch.ops.aten.addmm.default(self._bias, x, self._weight)
large_model_graph = torch.export(LargeModel(), inputs).graph
class PatternModel(torch.nn.Module):
def __init__(self):
super().__init__()
self._weight_1 = torch.nn.Parameter(torch.ones(5, 5))
self._bias_1 = torch.nn.Parameter(torch.ones(5, 5))
def forward(self, x):
return torch.ops.aten.addmm.default(self._bias_1, x, self._weight_1)
pattern_graph = torch.export(PatternModel(), inputs).graph
subgraph_matcher = SubgraphMatcher(pattern_graph)
match_result = subgraph_matcher.match(large_model_graph)
Функция match возвращает список InternalMatch:
@dataclass
class InternalMatch():
# Nodes from which the match was found
anchors: List[Node]
# Maps nodes in the pattern subgraph to nodes in the larger graph
nodes_map: Dict[Node, Node] = field(default_factory=dict)
# Nodes in target graph that are matched placeholder in pattern
placeholder_nodes: List[Node] = field(default_factory=list)
# Nodes in matched subgraph returned by output
returning_nodes: List[Node] = field(default_factory=list)
Разделитель на основе возможностей
Чтобы найти самые большие подграфы узлов, поддерживающих определённый инвариант, можно использовать CapabilityBasedPartitioner FX.
Атрибуты класса
-
graph_module (torch.fx.GraphModule): модуль графа, который мы разбиваем на разделы. -
operator_support (OperatorSupportBase): объект, определяющий, поддерживается ли узел графа в разделе. -
allows_single_node_partition (bool): если True, разрешает создавать разделы из одного узла. -
non_compute_ops (Optional[Sequence[str]]): набор операций, считающихся «невычислительными» (например,torch.ops.aten.viewи_operator.getitem), чтобы разделитель не создавал графы, содержащие только такие невычислительные операции -
allowed_single_node_partition_ops (Optional[Sequence[str]]): набор операций, которые разрешено включать в раздел из одного узла.
Класс OperatorSupportBase используется разделителем для определения того, относится ли определённый узел графа к разделу. Для этого переопределяется функция is_node_supported. Можно объединить несколько объектов OperatorSupportBase с помощью chain (возвращает False, если любой из объектов OperatorSupportBase возвращает False) и any_chain (возвращает True, если любой из объектов OperatorSupportBase возвращает True).
Пример:
from torch.fx.passes.infra.partitioner import CapabilityBasedPartitioner
from torch.fx.passes.operator_support import any_chain, OperatorSupportBase
class AddMulOperatorSupport(OperatorSupportBase):
def is_node_supported(self, submodules, node: torch.fx.Node) -> bool:
return node.op == "call_function" and node.target in [
torch.ops.aten.add.Tensor, torch.ops.aten.mul.Tensor,
]
capability_partitioner = CapabilityBasedPartitioner(
graph_module,
op_support,
)
# Returns a list of partitions (list of nodes that belong in each partition)
partition_list = capability_partitioner.propose_partitions()
# Fuses the partitions into graph modules and inserts `call_module` nodes in the graph
fused_graph_module = capability_partitioner.fuse_partitions(partition_list)
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/user_guide/torch_compiler/torch.compiler_transformations.html