Spec-Zone.ru › PyTorch 2

Написание преобразований графов в ATen IR

Пасы

Поскольку 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 уровня.

Преобразователь

Для случаев использования уровня 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 subgraph rewriter. Дана pattern, он создаёт подграф операторов, соответствующих шаблону, и затем заменяет каждый совпавший подграф на replacement.

Примечание:

This is an inplace operation.

Входные данные pattern и replacement должны быть вызываемыми функциями или GraphModules, содержащими те же операторы, которые используются в графе (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 <https://github.com/pytorch/pytorch/blob/main/torch/fx/passes/infra/pass_manager.py>`__ — это класс, используемый для выполнения нескольких пасов над заданным модулем графа. При инициализации экземпляра 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 на основе графов для разбиения графа.

Сопоставитель подграфов

Для поиска подграфов в графе, которые соответствуют определённому шаблону, мы можем использовать сопоставитель подграфов FX `SubgraphMatcher <https://github.com/pytorch/pytorch/blob/main/torch/fx/passes/utils/matcher_utils.py>`__.

Атрибуты класса:

  • 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)

Партишнер на основе возможностей

Для поиска самых больших подграфов узлов, поддерживающих определённое свойство, мы можем использовать партишнер на основе возможностей FX `CapabilityBasedPartitioner <https://github.com/pytorch/pytorch/blob/main/torch/fx/passes/infra/partitioner.py#L34>`__.

Атрибуты класса

  • 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 <https://github.com/pytorch/pytorch/blob/main/torch/fx/passes/operator_support.py#LL28C1-L28C1>`__ используется партишенером для определения, принадлежит ли конкретный узел в графе партиции. Это делается путём переопределения функции is_node_supported. Вы можете комбинировать несколько OperatorSuppportBase с помощью `chain <https://github.com/pytorch/pytorch/blob/main/torch/fx/passes/operator_support.py#L150>`__(возвращает False, если любой из OperatorSupportBase возвращает False) и `any_chain <https://github.com/pytorch/pytorch/blob/main/torch/fx/passes/operator_support.py#L164>`__ (возвращает 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)

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/torch.compiler_transformations.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API