Spec-Zone.ru › PyTorch 2.14

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. Это позволяет значительно автоматизировать утомительный код для манипуляций с графом, который может стать громоздким по мере усложнения преобразований.

Примеры манипуляций с графом

  • Замена одной операции
  • Объединение Conv/Batch Norm
  • replace_pattern: базовое использование
  • Квантование
  • Инвертирующее преобразование

Прокси и повторная трассировка

Еще один способ манипулировать 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, к которому применены правила преобразования, заданные переопределенными методами.

Примеры использования шаблона «Интерпретатор»

  • Распространение формы
  • Профилировщик производительности

Отладка

Введение

При разработке преобразований нередко оказывается, что код работает не совсем правильно. В таком случае может потребоваться отладка. Важно действовать в обратном порядке: сначала проверьте результаты вызова сгенерированного модуля, чтобы подтвердить или опровергнуть его корректность. Затем изучите и отладьте сгенерированный код. После этого отладьте процесс преобразований, который привел к созданию этого кода.

Если вы не знакомы с отладчиками, см. вспомогательный раздел Доступные отладчики.

Типичные ошибки при создании преобразований

  • Недетерминированный порядок перебора set. В Python тип данных set неупорядочен. Использование set для хранения коллекций объектов, например Node\ s, может привести к неожиданной недетерминированности. Например, можно перебирать набор Node s и добавлять их в Graph. Поскольку тип данных set неупорядочен, порядок операций в выходной программе будет недетерминированным и может меняться при каждом запуске программы. Рекомендуется использовать тип данных dict, который сохраняет порядок добавления элементов начиная с Python 3.7 (и cPython 3.6). dict можно использовать как множество, сохраняя значения для удаления дубликатов в ключах dict.

Проверка корректности модулей

Поскольку выходные данные большинства модулей глубокого обучения представляют собой экземпляры torch.Tensor с плавающей точкой, проверить эквивалентность результатов двух torch.nn.Module не так просто, как выполнить обычное сравнение на равенство. Для пояснения рассмотрим пример:

import torch
import torch.fx
import torchvision.models as models

def transform(m : torch.nn.Module) -> torch.nn.Module:
    gm = torch.fx.symbolic_trace(m)

    # Imagine we're doing some transforms here
    # <...>

    gm.recompile()

    return gm

resnet18 = models.resnet18()
transformed_resnet18 = transform(resnet18)

input_image = torch.randn(5, 3, 224, 224)

assert resnet18(input_image) == transformed_resnet18(input_image)
"""
RuntimeError: Boolean value of Tensor with more than one value is ambiguous
"""

Здесь мы попытались проверить равенство значений двух моделей глубокого обучения с помощью оператора равенства ==. Однако такое сравнение некорректно: оператор возвращает тензор, а не 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 в настоящее время не поддерживаются.
    • Аннотации локальных имён внутри функции в настоящее время не поддерживаются.
  • Особенность флага 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)
    
  • Из-за этого различия рассмотрите возможность пометить модули, динамически взаимодействующие с флагом 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 * 2

FX обычно не может трассировать этот код из-за наличия потока управления. Однако с помощью 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.

Тип возвращаемого значения:

GraphModule

Примечание

Обратная совместимость этого 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 или имя, подходящее для контекста вашего преобразования.

Примечание

Обратная совместимость этого API гарантируется.

add_submodule(target, m) [исходный код]

Добавляет указанный подмодуль в self.

Если ещё не существуют промежуточные модули, являющиеся частью пути target, для них создаются пустые Module.

Параметры:
  • target (str) – Полное строковое имя нового подмодуля (пример задания полного строкового имени см. в nn.Module.get_submodule)
  • m (Module) – Сам подмодуль; фактический объект, который нужно установить в текущий Module
Возвращает:
Удалось ли вставить подмодуль. Чтобы

этот метод вернул True, каждый объект в цепочке, обозначенной target, должен либо а) ещё не существовать, либо б) ссылаться на nn.Module (а не на параметр или другой атрибут)

Тип возвращаемого значения:

bool

Примечание

Обратная совместимость этого 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 не является допустимой ссылкой на подмодуль.

Тип возвращаемого значения:

bool

Примечание

Обратная совместимость этого 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”]).

Тип возвращаемого значения:

str

Предупреждение

Этот 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_function Node в 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
Возвращает:

Недавно созданный и вставленный узел call_function.

Тип возвращаемого значения:

Node

Примечание

Для этого метода действуют те же правила точки вставки и выражения типа, что и для Graph.create_node().

Примечание

Для этого API гарантирована обратная совместимость.

call_method(method_name, args=None, kwargs=None, type_expr=None) [исходный код]

Вставить call_method Node в 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, который будет иметь результат этого узла.
Возвращает:

Недавно созданный и вставленный узел call_method.

Тип возвращаемого значения:

Node

Примечание

Для этого метода действуют те же правила точки вставки и выражения типа, что и для Graph.create_node().

Примечание

Для этого API гарантирована обратная совместимость.

call_module(module_name, args=None, kwargs=None, type_expr=None) [исходный код]

Вставить call_module Node в 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, который будет иметь результат этого узла.
Возвращает:

Недавно созданный и вставленный узел call_module.

Тип возвращаемого значения:

Node

Примечание

Для этого метода действуют те же правила точки вставки и выражения типа, что и для 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, который будет иметь результат этого узла.
Возвращает:

Недавно созданный и вставленный узел.

Тип возвращаемого значения:

Node

Примечание

Для этого API гарантирована обратная совместимость.

create_size_node(tensor_node, dim) [исходный код]

Создать узел FX для tensor_node.size(dim).

Предупреждение

Этот API является экспериментальным и НЕ обеспечивает обратную совместимость.

Тип возвращаемого значения:

Node

create_storage_offset_node(tensor_node) [исходный код]

Создать узел FX для tensor_node.storage_offset().

Предупреждение

Этот API является экспериментальным и НЕ обеспечивает обратную совместимость.

Тип возвращаемого значения:

Node

create_stride_node(tensor_node, dim) [исходный код]

Создать узел FX для tensor_node.stride(dim).

Предупреждение

Этот API является экспериментальным и НЕ обеспечивает обратную совместимость.

Тип возвращаемого значения:

Node

eliminate_dead_code(is_impure_node=None) [исходный код]

Удалить из графа весь неиспользуемый код, учитывая количество пользователей каждого узла и наличие у узлов побочных эффектов. Перед вызовом граф должен быть отсортирован топологически.

Параметры:
  • is_impure_node (Optional[Callable[[Node], bool]]) – Функция, возвращающая
  • None (является ли узел нечистым. Если это) –
  • to (то поведение по умолчанию —) –
  • Node.is_impure. (использовать) –
Возвращает:

Был ли граф изменён в результате этого прохода.

Тип возвращаемого значения:

bool

Пример:

До удаления неиспользуемого кода у 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) [исходный код]

Позволяет быстро выполнять поиск узлов

Параметры:
  • op (str) – имя операции
  • target (Optional[Target]) – цель узла. Для call_function цель обязательна. Для других операций цель указывать необязательно.
  • sort (bool) – нужно ли возвращать узлы в том порядке, в котором они встречаются в графе.
Возвращает:

Итерируемый объект с узлами, соответствующими заданным операции и цели.

Тип возвращаемого значения:

list[Any]

Предупреждение

Этот API является экспериментальным и НЕ обеспечивает обратную совместимость.

get_attr(qualified_name, type_expr=None) [исходный код]

Вставить узел get_attr в граф. Узел get_attr Node представляет получение атрибута из иерархии Module.

Параметры:
  • qualified_name (str) – полное квалифицированное имя атрибута, который требуется получить. Например, если трассированный Module содержит подмодуль с именем foo, у которого есть подмодуль с именем bar, содержащий атрибут с именем baz, в качестве qualified_name следует передать полное квалифицированное имя foo.bar.baz.
  • type_expr (Optional[Any]) – необязательная аннотация типа, представляющая тип Python, который будет иметь результат этого узла.
Возвращает:

Недавно созданный и вставленный узел get_attr.

Тип возвращаемого значения:

Node

Примечание

Для этого метода действуют те же правила точки вставки и выражения типа, что и для Graph.create_node.

Примечание

Для этого API гарантирована обратная совместимость.

graph_copy(g, val_map, return_output_node=False) [исходный код]

Скопировать все узлы из заданного графа в self.

Параметры:
  • g (Graph) – Исходный граф, из которого копируются узлы.
  • val_map (Dict[Node, Node]) – словарь, который будет заполнен сопоставлениями узлов из g узлам из self. Обратите внимание: val_map можно передать уже заполненным, чтобы переопределить копирование определённых значений.
Возвращает:

Значение в 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 является экспериментальным и НЕ обеспечивает обратную совместимость.

Тип возвращаемого значения:

Node | int

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)]) обходит символьное выражение SymPy s32 и создаёт подграф 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 является экспериментальным и НЕ гарантирует обратную совместимость.

Тип возвращаемого значения:

list[Node | int]

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])
Параметры:
  • node (Node) – Узел, который нужно скопировать в self.
  • arg_transform (Callable[[Node], Argument]) – Функция, преобразующая аргументы Node в args и kwargs узла в эквивалентный аргумент в self. В простейшем случае она должна получать значение из таблицы, сопоставляющей узлы исходного графа с self.
Тип возвращаемого значения:

Node

Примечание

Для этого 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) [исходный код]

Вставляет узел output Node в Graph. Узел output представляет инструкцию return в коде Python. result — это возвращаемое значение.

Параметры:
  • result (Argument) – Возвращаемое значение.
  • type_expr (Optional[Any]) – необязательная аннотация типа, представляющая тип Python, который будет иметь выходное значение этого узла.

Примечание

К этому методу применяются те же правила для позиции вставки и выражения типа, что и к Graph.create_node.

Примечание

Для этого API гарантирована обратная совместимость.

output_node() [исходный код]

Предупреждение

Этот API является экспериментальным и НЕ гарантирует обратную совместимость.

Тип возвращаемого значения:

Node

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, указывая, что параметр _не_ имеет значения по умолчанию.
Тип возвращаемого значения:

Node

Примечание

К этому методу применяются те же правила для позиции вставки и выражения типа, что и к Graph.create_node.

Примечание

Для этого API гарантирована обратная совместимость.

print_tabular() [исходный код]

Выводит промежуточное представление графа в табличном формате. Обратите внимание, что для этого API требуется установленный модуль tabulate.

Примечание

Для этого API гарантирована обратная совместимость.

process_inputs(*args) [исходный код]

Обрабатывает аргументы, чтобы их можно было передать графу FX.

Предупреждение

Этот API является экспериментальным и НЕ гарантирует обратную совместимость.

Тип возвращаемого значения:

Any

process_outputs(out) [исходный код]

Предупреждение

Этот API является экспериментальным и НЕ гарантирует обратную совместимость.

Тип возвращаемого значения:

Any

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) – Нужно ли включать метаданные тензора
Возвращает:
If 1) we’re using format_node as an internal helper

в методе __str__ класса Graph, и 2) если self является узлом-заполнителем, возвращает None. В противном случае возвращает описательное строковое представление текущего узла.

Тип возвращаемого значения:

str

Примечание

Обратная совместимость этого API гарантируется.

insert_arg(idx, arg) [исходный код]

Вставляет позиционный аргумент в список аргументов по заданному индексу.

Параметры:
  • idx (int) – Индекс элемента в self.args, перед которым нужно выполнить вставку.
  • arg (Argument) – Новое значение аргумента, которое нужно вставить в args

Примечание

Обратная совместимость этого API гарантируется.

is_impure(impure_random=True) [исходный код]

Возвращает, является ли эта операция нечистой, то есть является ли её тип placeholder или output либо call_function или call_module, если такая операция нечистая.

Параметры:

impure_random (bool) – Следует ли считать операцию rand нечистой.

Возвращает:

Является ли операция нечистой.

Тип возвращаемого значения:

bool

Предупреждение

Этот 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.
Возвращает:

Список узлов, в которых было выполнено это изменение.

Тип возвращаемого значения:

list[Node]

Примечание

Обратная совместимость этого API гарантируется.

replace_input_with(old_input, new_input) [исходный код]

Перебирает входные узлы self и заменяет все вхождения old_input на new_input.

Параметры:
  • old_input (Node) – Старый входной узел, который нужно заменить.
  • new_input (Node) – Новый входной узел, заменяющий old_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

Примечание

Обратная совместимость этого API гарантируется.

update_kwarg(key, arg) [исходный код]

Обновляет существующий именованный аргумент, задавая ему новое значение arg. После вызова self.kwargs[key] == arg.

Параметры:
  • key (str) – Ключ обновляемого элемента в self.kwargs
  • arg (Argument) – Новое значение аргумента для записи в kwargs

Примечание

Обратная совместимость этого 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.

Тип возвращаемого значения:

Any

Примечание

Для этого API гарантируется обратная совместимость.

create_arg(a) [исходный код]

Метод, определяющий поведение трассировки при подготовке значений, используемых в качестве аргументов узлов в Graph.

По умолчанию выполняются следующие действия:

  1. Перебираются коллекции (например, tuple, list, dict), а для их элементов рекурсивно вызывается create_args.
  2. Если передан объект Proxy, возвращается ссылка на базовый IR Node
  3. Если передан объект Tensor, не являющийся Proxy, для разных случаев создаётся IR:

    • Для Parameter создаётся узел get_attr, ссылающийся на этот Parameter
    • Tensor, не являющийся Parameter, сохраняется в специальном атрибуте, который ссылается на этот атрибут.

Этот метод можно переопределить для поддержки дополнительных типов.

Параметры:

a (Any) – Значение, которое будет создано как Argument в Graph.

Возвращает:

Значение a, преобразованное в соответствующий Argument

Тип возвращаемого значения:

Argument

Примечание

Для этого API гарантируется обратная совместимость.

create_args_for_root(root_fn, is_module, concrete_args=None) [исходный код]

Создаёт узлы placeholder, соответствующие сигнатуре корневого модуля root. Этот метод анализирует сигнатуру root и создаёт соответствующие узлы, поддерживая также *args и **kwargs.

Предупреждение

Этот API является экспериментальным и НЕ имеет обратной совместимости.

Тип возвращаемого значения:

tuple[Any, list[Any]]

create_node(kind, target, args, kwargs, name=None, type_expr=None) [исходный код]

Вставляет узел графа с заданными target, args, kwargs и name.

Этот метод можно переопределить для дополнительной проверки, валидации или изменения значений, используемых при создании узла. Например, можно запретить запись операций на месте.

Примечание

Для этого API гарантируется обратная совместимость.

Тип возвращаемого значения:

Node

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 гарантируется обратная совместимость.

Тип возвращаемого значения:

Proxy

get_fresh_qualname(prefix) [исходный код]

Получает новое имя для префикса и возвращает его. Эта функция гарантирует, что оно не совпадёт с уже существующим атрибутом графа.

Примечание

Для этого API гарантируется обратная совместимость.

Тип возвращаемого значения:

str

getattr(attr, attr_val, parameter_proxy_cache) [исходный код]

Метод, определяющий поведение этого Tracer при вызове getattr для вызова экземпляра nn.Module.

По умолчанию возвращается прокси-значение атрибута. Оно также сохраняется в parameter_proxy_cache, чтобы при последующих вызовах использовать существующий proxy, а не создавать новый.

Этот метод можно переопределить, например, чтобы не возвращать прокси при запросе параметров.

Параметры:
  • attr (str) – Имя запрашиваемого атрибута
  • attr_val (Any) – Значение атрибута
  • parameter_proxy_cache (Dict[str, Any]) – Кэш имён атрибутов и соответствующих прокси
Возвращает:

Значение, возвращённое вызовом getattr.

Тип возвращаемого значения:

Any

Предупреждение

Этот API является экспериментальным и НЕ имеет обратной совместимости.

is_leaf_module(m, module_qualified_name) [исходный код]

Метод, определяющий, является ли заданный nn.Module «листовым» модулем.

Листовые модули — это атомарные единицы, которые появляются в IR и на которые ссылаются вызовы call_module. По умолчанию модули из пространства имён стандартной библиотеки PyTorch (torch.nn) являются листовыми. Все остальные модули трассируются, а их составляющие операции записываются, если для них не задано иное значение этого параметра.

Параметры:
  • m (Module) – Проверяемый модуль
  • module_qualified_name (str) – Путь к корню этого модуля. Например, если в иерархии модулей подмодуль foo содержит подмодуль bar, который, в свою очередь, содержит подмодуль baz, то здесь этот модуль будет представлен полным именем foo.bar.baz.
Тип возвращаемого значения:

bool

Примечание

Для этого API гарантируется обратная совместимость.

iter(obj) [исходный код]
Вызывается при переборе объекта proxy, например

при использовании в потоке управления. Обычно неизвестно, что делать, поскольку значение proxy неизвестно, однако пользовательский трассировщик может добавить к узлу графа дополнительные сведения с помощью create_node и вернуть итератор.

Примечание

Для этого API гарантируется обратная совместимость.

Тип возвращаемого значения:

Iterator

keys(obj) [исходный код]
Вызывается при вызове метода keys() у объекта proxy.

Так происходит при вызове ** для proxy. Если ** должен работать в пользовательском трассировщике, этот метод должен возвращать итератор.

Примечание

Для этого API гарантируется обратная совместимость.

Тип возвращаемого значения:

Proxy

path_of_module(mod) [исходный код]

Вспомогательный метод для определения полного имени mod в иерархии модулей root. Например, если у root есть подмодуль с именем foo, у которого есть подмодуль с именем bar, то при передаче bar в эту функцию будет возвращена строка «foo.bar».

Параметры:

mod (str) – Module, для которого нужно получить полное имя.

Тип возвращаемого значения:

str

Примечание

Для этого API гарантируется обратная совместимость.

proxy(node) [исходный код]

Примечание

Для этого API гарантируется обратная совместимость.

Тип возвращаемого значения:

Proxy

to_bool(obj) [исходный код]
Вызывается при преобразовании объекта proxy в логическое значение, например

при использовании в потоке управления. Обычно неизвестно, что делать, поскольку значение proxy неизвестно, однако пользовательский трассировщик может добавить к узлу графа дополнительные сведения с помощью create_node и вернуть значение.

Примечание

Для этого API гарантируется обратная совместимость.

Тип возвращаемого значения:

bool

trace(root, concrete_args=None) [исходный код]

Трассирует root и возвращает соответствующее представление FX Graph. root может быть экземпляром nn.Module или вызываемым объектом Python.

Обратите внимание: после этого вызова self.root может отличаться от переданного сюда root. Например, если в trace() передана свободная функция, для использования в качестве корневого объекта и добавления встроенных констант будет создан экземпляр nn.Module.

Параметры:
  • root (Union[Module, Callable]) – Объект Module или функция, которую нужно трассировать. Для этого параметра гарантируется обратная совместимость.
  • concrete_args (Optional[Dict[str, any]]) – Конкретные аргументы, которые не следует считать Proxy. Этот параметр является экспериментальным, и его обратная совместимость НЕ гарантируется.
Возвращает:

Объект Graph, представляющий семантику переданного root.

Тип возвращаемого значения:

Graph

Примечание

Для этого 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 гарантируется.

Тип возвращаемого значения:

Any

call_function(target, args, kwargs) [исходный код]

Выполняет узел call_function и возвращает результат.

Параметры:
  • target (Target) – Цель вызова для этого узла. Подробности семантики см. в описании Node
  • args (Tuple) – Кортеж позиционных аргументов этого вызова
  • kwargs (Dict) – Словарь именованных аргументов этого вызова
Тип возвращаемого значения:

Any

Возвращает

Any: значение, возвращённое при вызове функции

Примечание

Обратная совместимость этого API гарантируется.

call_method(target, args, kwargs) [исходный код]

Выполняет узел call_method и возвращает результат.

Параметры:
  • target (Target) – Цель вызова для этого узла. Подробности семантики см. в описании Node
  • args (Tuple) – Кортеж позиционных аргументов этого вызова
  • kwargs (Dict) – Словарь именованных аргументов этого вызова
Тип возвращаемого значения:

Any

Возвращает

Any: значение, возвращённое при вызове метода

Примечание

Обратная совместимость этого API гарантируется.

call_module(target, args, kwargs) [исходный код]

Выполняет узел call_module и возвращает результат.

Параметры:
  • target (Target) – Цель вызова для этого узла. Подробности семантики см. в описании Node
  • args (Tuple) – Кортеж позиционных аргументов этого вызова
  • kwargs (Dict) – Словарь именованных аргументов этого вызова
Тип возвращаемого значения:

Any

Возвращает

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 гарантируется.

Тип возвращаемого значения:

Any

call_module(target, args, kwargs) [исходный код]

Примечание

Обратная совместимость этого API гарантируется.

Тип возвращаемого значения:

Any

get_attr(target, args, kwargs) [исходный код]

Выполняет узел get_attr. В Transformer этот метод переопределён для вставки нового узла get_attr в выходной граф.

Параметры:
  • target (Target) – Цель вызова для этого узла. Подробности семантики см. в описании Node
  • args (Tuple) – Кортеж позиционных аргументов этого вызова
  • kwargs (Dict) – Словарь именованных аргументов этого вызова
Тип возвращаемого значения:

Proxy

Примечание

Обратная совместимость этого API гарантируется.

placeholder(target, args, kwargs) [исходный код]

Выполняет узел placeholder. В Transformer этот метод переопределён для вставки нового узла placeholder в выходной граф.

Параметры:
  • target (Target) – Цель вызова для этого узла. Подробности семантики см. в описании Node
  • args (Tuple) – Кортеж позиционных аргументов этого вызова
  • kwargs (Dict) – Словарь именованных аргументов этого вызова
Тип возвращаемого значения:

Proxy

Примечание

Обратная совместимость этого API гарантируется.

transform() [исходный код]

Преобразует self.module и возвращает преобразованный GraphModule.

Примечание

Обратная совместимость этого API гарантируется.

Тип возвращаемого значения:

GraphModule

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, а параметры Callable replacement должны соответствовать шаблону. Первое правило объясняет, почему в приведённом выше блоке кода функция 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) – Модуль графа для топологической сортировки. Он изменяется на месте.

Возвращает:

Модуль графа, отсортированный на месте

Тип возвращаемого значения:

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

Для оператора из пространства имён torch возвращает список объектов inspect.Signature, соответствующих перегрузкам этого оператора.

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

Находит все возможные непересекающиеся наборы операторов и их зависимостей по данным (pattern) в графе GraphModule (gm), а затем заменяет каждый найденный подграф другим подграфом (replacement).

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

Задаёт частичный порядок (функцию «зависит от»), в котором this должен выполняться до that.

torch.fx.passes.operator_support

any_chain

Объединяет последовательность экземпляров OperatorSupportBase в один экземпляр OperatorSupportBase

chain

Объединяет последовательность экземпляров OperatorSupportBase в один экземпляр OperatorSupportBase

create_op_support

Оборачивает функцию IsNodeSupported в экземпляр OperatorSupportBase

torch.fx.passes.param_fetch

torch.fx.passes.param_fetch.default_matching(name, target_version) [исходный код]

Метод сопоставления по умолчанию

Предупреждение

Этот API является экспериментальным, и его обратная совместимость НЕ гарантируется.

Тип возвращаемого значения:

str

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 является экспериментальным, и его обратная совместимость НЕ гарантируется.

Тип возвращаемого значения:

dict[str, Any]

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

Задаёт частичный порядок (функцию «зависит от»), в котором these должны выполняться до those.

this_before_that_pass_constraint

Задаёт частичный порядок (функцию «зависит от»), в котором this должен выполняться до that.

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 linear1

main:

def forward(self, in2, linear1):
    self = self.root
    linear2 = self.linear2(in2)
    cat_1 = torch.cat([linear1, linear2])
    linear3 = self.linear3(cat_1)
    return linear3

main:

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 является экспериментальным, и его обратная совместимость НЕ гарантируется.

setattr_recursive

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 в gm

original_inputs (Tuple[Node, …]): входные узлы для nodes в исходном gm

original_outputs (Tuple[Node, …]): узлы-потребители nodes в исходном gm

Тип возвращаемого значения:

fused_gm (GraphModule)

Предупреждение

Этот API является экспериментальным и НЕ обладает обратной совместимостью.

erase_nodes
topo_sort
validate_partition

torch.fx.passes.utils.source_matcher_utils

torch.fx.passes.utils.source_matcher_utils.get_source_partitions(graph, wanted_sources, filter_fn=None) [исходный код]
Параметры:
  • graph (Graph) – Граф, который требуется разбить на части
  • wanted_sources (list[Any]) – Список источников узлов, полученных декомпозицией из этого источника. Это может быть функция (например, torch.nn.functional.linear) или тип листового модуля (например, torch.nn.Linear).
Возвращает:

Словарь, сопоставляющий переданные источники со списком SourcePartition, соответствующих списку узлов, декомпозированных из указанного источника.

Тип возвращаемого значения:

dict[Any, list[SourcePartition]]

Предупреждение

Этот API является экспериментальным и НЕ обладает обратной совместимостью.

check_subgraphs_connected

Для двух подграфов A и B (представленных в виде списков узлов) проверяет, есть ли в A узлы, соединённые хотя бы с одним узлом в B, то есть существует ли узел в B, который использует узел из A (но не наоборот).

get_unique_attr_name_in_module

Проверяет, что имя уникально (в модуле) и может представлять атрибут.

split_const_subgraphs

Просматривает module в поисках узлов, у которых все входные атрибуты являются константами, и выделяет их в отдельный константный подграф. Возвращает FoldedGraphModule, который при первом запуске выполняет этот константный подграф, чтобы задать атрибуты модуля перед выполнением неконстантной части графа.

CodeGen
reduce_graph_module
reduce_package_graph_module
torch.fx.passes.annotate_getitem_nodes.annotate_getitem_nodes(graph) [исходный код]

Аннотирует тип узлов getitem, определяемый по типу узла последовательности. Если узел последовательности не аннотирован типом, ничего не делает. В настоящее время поддерживает узлы getitem для узлов последовательностей tuple, list и NamedTuple.

Это полезно, поскольку аннотации локальных имён внутри функции теряются при преобразованиях FX. Добавление известных аннотаций типов обратно к узлам getitem повышает совместимость со скриптами JIT.

Параметры:

graph (torch.fx.Graph) – Граф, который нужно аннотировать

GraphTransformObserver
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) [исходный код]

Перемещает узлы, не являющиеся тензорами, на границу между подграфами.

Для каждого подграфа:

  1. Находит узлы, тип которых не является тензором и у которых есть потомки в другом подграфе, и помещает их в очередь для следующего шага
  2. Выполняет BFS для узлов в очереди и DFS для каждого узла; допустим, это узел X, находящийся в подграфе A:

    1. если он находится в to_subgraph, возвращает результат (продолжает DFS)
    2. если он находится в from_subgraph, добавляет узлы в nodes_to_move и продолжает DFS
    3. в противном случае это означает, что его нельзя переместить
    4. также проверяет, нужно ли добавить родительский узел X в очередь. (В очереди могут быть повторяющиеся узлы; каждый узел обрабатывается только один раз)
Параметры:

subgraphs (list[Subgraph]) – Список подграфов, содержащих узлы для обработки

Предупреждение

Этот API является экспериментальным и НЕ обладает обратной совместимостью.

torch.fx.passes.splitter_base.generate_inputs_for_submodules(model, inputs, target_submodules, deepcopy=False) [исходный код]

Формирует входные данные для целевых подмодулей заданной модели. Обратите внимание: эта функция не работает, если два подмодуля ссылаются на один и тот же объект.

Параметры:
  • model (Module) – корневая модель.
  • inputs (Sequence[Any]) – входные данные корневой модели.
  • target_submodules (Iterable[str]) – подмодули, для которых нужно сформировать входные данные.
Возвращает:

Словарь, сопоставляющий имена подмодулей с их входными данными.

Тип возвращаемого значения:

dict[str, Any]

Предупреждение

Этот API является экспериментальным и НЕ обладает обратной совместимостью.

NodeEvent

Событие, произошедшее с узлом при разделении графа.

NodeEventTracker

Отслеживает события узлов во время выполнения разделителя.

SubgraphMatcherWithNameNodeMap

Расширяет SubgraphMatcher, добавляя поддержку поиска узлов сопоставленного подграфа по имени узла,

GraphAppendingTracer

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/fx.html

Spec-Zone.ru

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