Spec-Zone.ru › PyTorch 2

torch.fx

Обзор

FX — это набор инструментов для разработчиков, предназначенных для преобразования nn.Module экземпляров. FX состоит из трёх основных компонентов: символического трассировщика, промежуточного представления и генерации Python-кода. Пример работы этих компонентов:

import torch
# Simple module for demonstration
class MyModule(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.param = torch.nn.Parameter(torch.rand(3, 4))
        self.linear = torch.nn.Linear(4, 5)

    def forward(self, x):
        return self.linear(x + self.param).clamp(min=0.0, max=1.0)

module = MyModule()

from torch.fx import symbolic_trace
# Symbolic tracing frontend - captures the semantics of the module
symbolic_traced : torch.fx.GraphModule = symbolic_trace(module)

# High-level intermediate representation (IR) - Graph representation
print(symbolic_traced.graph)
"""
graph():
    %x : [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-кода. Он подаёт фиктивные значения, называемые Прокси, через код. Операции над этими Прокси записываются. Дополнительную информацию о символическом трассировании можно найти в документации symbolic_trace() и Tracer.

Промежуточное представление — это контейнер для операций, записанных во время символического трассирования. Оно состоит из списка узлов, которые представляют входные данные функции, вызов функций, методов или экземпляров torch.nn.Module, и возвращаемые значения. Дополнительную информацию об IR можно найти в документации Graph. IR — это формат, на котором применяются преобразования.

Генерация Python-кода — то, что делает FX инструментом преобразования Python в Python (или Модуль в Модуль). Для каждой IR-структуры Graph мы можем сгенерировать допустимый Python-код, соответствующий семантике Graph. Эта функциональность обернута в GraphModule, который является экземпляром torch.nn.Module, содержащим Graph и forward метод, сгенерированный из Graph.

Вместе эта цепочка компонентов (символическое трассирование -> промежуточное представление -> преобразования -> генерация Python-кода) составляет цепочку преобразования Python в Python с помощью FX. Кроме того, эти компоненты могут использоваться по отдельности. Например, символическое трассирование может использоваться изолированно для захвата формы кода в целях анализа (а не преобразования). Генерация кода может использоваться для программно сгенерированных моделей, например, из файла конфигурации. У FX много применений!

Несколько примеров преобразований можно найти в репозитории примеров.

Написание Преобразований

Что такое преобразование FX? По сути, это функция, которая выглядит так.

import torch
import torch.fx

def transform(m: nn.Module,
              tracer_class : type = torch.fx.Tracer) -> torch.nn.Module:
    # Step 1: Acquire a Graph representing the code in `m`

    # NOTE: torch.fx.symbolic_trace is a wrapper around a call to
    # fx.Tracer.trace and constructing a GraphModule. We'll
    # split that out in our transform to allow the caller to
    # customize tracing behavior.
    graph : torch.fx.Graph = tracer_class().trace(m)

    # Step 2: Modify this Graph or create a new one
    graph = ...

    # Step 3: Construct a Module to return
    return torch.fx.GraphModule(m, graph)

Ваше преобразование будет принимать torch.nn.Module, получать Graph из него, выполнять некоторые изменения и возвращать новый torch.nn.Module. Вы должны рассматривать torch.nn.Module, возвращаемый вашим преобразованием FX, как идентичный обычному torch.nn.Module — вы можете передавать его другому преобразованию FX, TorchScript или запускать его. Гарантия того, что вход и выход вашего преобразования FX — это torch.nn.Module, обеспечит композиционность.

Примечание

Также можно изменить существующий GraphModule вместо создания нового, например:

import torch
import torch.fx

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

    # Modify gm.graph
    # <...>

    # Recompile the forward() method of `gm` from its Graph
    gm.recompile()

    return gm

Обратите внимание, что вы ДОЛЖНЫ вызвать GraphModule.recompile(), чтобы синхронизировать сгенерированный forward() метод на GraphModule с изменённой Graph.

Учитывая, что вы передали torch.nn.Module, который был прослежен в Graph, у вас есть два основных подхода к созданию нового Graph.

Краткий обзор графов

Полное описание семантики графов можно найти в документации Graph, но мы рассмотрим основы здесь. Graph — это структура данных, которая представляет метод GraphModule. Для этого необходима информация:

  • Какие входные данные метода?
  • Какие операции выполняются внутри метода?
  • Какое значение возвращается (т. е. возвращаемое значение) методом?

Все три этих понятия представлены экземплярами Node. Посмотрим на это на коротком примере:

import torch
import torch.fx

class MyModule(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.param = torch.nn.Parameter(torch.rand(3, 4))
        self.linear = torch.nn.Linear(4, 5)

    def forward(self, x):
        return torch.topk(torch.sum(
            self.linear(x + self.linear.weight).relu(), dim=-1), 3)

m = MyModule()
gm = torch.fx.symbolic_trace(m)

gm.graph.print_tabular()

Здесь мы определяем модуль MyModule для демонстрационных целей, инициализируем его, символично отслеживаем его, а затем вызываем метод Graph.print_tabular() для вывода таблицы, показывающей узлы этого Graph:

opcode

name

target

args

kwargs

placeholder

x

x

()

{}

get_attr

linear_weight

linear.weight

()

{}

call_function

add_1

<built-in function add>

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

Для простых преобразований, состоящих только из подстановок, вы также можете использовать переписыватель подграфов subgraph rewriter.

Переписывание подграфов с помощью replace_pattern()

FX также предоставляет ещё один уровень автоматизации поверх непосредственного управления графом. API replace_pattern() по сути является инструментом «найти/заменить» для редактирования графов Graph. Он позволяет указать функцию pattern и функцию replacement и будет отслеживать эти функции, искать экземпляры группы операций в графе pattern и заменять эти экземпляры копиями графа replacement. Это может помочь значительно автоматизировать рутинную работу с графами, которая может стать неуправляемой по мере усложнения преобразований.

Примеры манипулирования графом

  • Замена одной операции
  • Слияние Conv/Batch Norm
  • replace_pattern: Базовое использование
  • Квантование
  • Инвертирование преобразования

Прокси/Переотслеживание

Другой способ манипулирования графами Graph — это повторное использование механизма Proxy, используемого при символическом отслеживании. Например, предположим, что мы хотим написать преобразование, которое разбивает функции PyTorch на более мелкие операции. Оно бы преобразовало каждый вызов F.relu(x) в (x > 0) * x. Одним из вариантов было бы выполнить необходимое переписывание графа, чтобы вставить сравнение и умножение после F.relu, а затем очистить исходный F.relu. Однако мы можем автоматизировать этот процесс, используя объекты Proxy, чтобы автоматически записывать операции в граф Graph.

Для использования этого метода мы пишем операции, которые мы хотим вставить, как обычный код PyTorch, и вызываем этот код с объектами Proxy в качестве аргументов. Эти объекты Proxy будут захватывать операции, выполняемые над ними, и добавлять их в граф Graph.

# Note that this decomposition rule can be read as regular Python
def relu_decomposition(x):
    return (x > 0) * x

decomposition_rules = {}
decomposition_rules[F.relu] = relu_decomposition

def decompose(model: torch.nn.Module,
              tracer_class : type = fx.Tracer) -> torch.nn.Module:
    """
    Decompose `model` into smaller constituent operations.
    Currently,this only supports decomposing ReLU into its
    mathematical definition: (x > 0) * x
    """
    graph : fx.Graph = tracer_class().trace(model)
    new_graph = fx.Graph()
    env = {}
    tracer = torch.fx.proxy.GraphAppendingTracer(new_graph)
    for node in graph.nodes:
        if node.op == 'call_function' and node.target in decomposition_rules:
            # By wrapping the arguments with proxies,
            # we can dispatch to the appropriate
            # decomposition rule and implicitly add it
            # to the Graph by symbolically tracing it.
            proxy_args = [
                fx.Proxy(env[x.name], tracer) if isinstance(x, fx.Node) else x for x in node.args]
            output_proxy = decomposition_rules[node.target](*proxy_args)

            # Operations on `Proxy` always yield new `Proxy`s, and the
            # return value of our decomposition rule is no exception.
            # We need to extract the underlying `Node` from the `Proxy`
            # to use it in subsequent iterations of this transform.
            new_node = output_proxy.node
            env[node.name] = new_node
        else:
            # Default case: we don't have a decomposition rule for this
            # node, so just copy the node over into the new graph.
            new_node = new_graph.node_copy(node, lambda x: env[x.name])
            env[node.name] = new_node
    return fx.GraphModule(model, new_graph)

Помимо избежания явного манипулирования графом, использование объектов Proxy также позволяет определять правила переписывания как обычный код Python. Для преобразований, требующих большого количества правил переписывания (таких как vmap или grad), это часто повышает читаемость и поддерживаемость правил. Обратите внимание, что при вызове Proxy мы также передавали трейсер, указывающий на базовую переменную graph. Это делается в случае, если операции в графе являются n-арными (например, add — бинарный оператор), вызов Proxy не создает несколько экземпляров трейсера графа, что может привести к неожиданным ошибкам во время выполнения. Мы рекомендуем этот метод использования объектов Proxy, особенно когда невозможно безопасно предположить, что базовые операторы являются унарными.

Пример практического использования объектов Proxy для манипулирования графами Graph можно найти здесь.

Шаблон интерпретатора

Полезным шаблоном организации кода в FX является перебор всех узлов Node в графе Graph и выполнение их. Это можно использовать для нескольких целей, включая аналитику значений, протекающих через граф во время выполнения, или преобразование кода посредством переотслеживания с использованием объектов Proxy. Например, предположим, что мы хотим выполнить GraphModule и записывать свойства формы и типа данных тензора torch.Tensor для узлов во время выполнения. Это может выглядеть так:

import torch
import torch.fx
from torch.fx.node import Node

from typing import Dict

class ShapeProp:
    """
    Shape propagation. This class takes a `GraphModule`.
    Then, its `propagate` method executes the `GraphModule`
    node-by-node with the given arguments. As each operation
    executes, the ShapeProp class stores away the shape and
    element type for the output values of each operation on
    the `shape` and `dtype` attributes of the operation's
    `Node`.
    """
    def __init__(self, mod):
        self.mod = mod
        self.graph = mod.graph
        self.modules = dict(self.mod.named_modules())

    def propagate(self, *args):
        args_iter = iter(args)
        env : Dict[str, Node] = {}

        def load_arg(a):
            return torch.fx.graph.map_arg(a, lambda n: env[n.name])

        def fetch_attr(target : str):
            target_atoms = target.split('.')
            attr_itr = self.mod
            for i, atom in enumerate(target_atoms):
                if not hasattr(attr_itr, atom):
                    raise RuntimeError(f"Node referenced nonexistant target {'.'.join(target_atoms[:i])}")
                attr_itr = getattr(attr_itr, atom)
            return attr_itr

        for node in self.graph.nodes:
            if node.op == 'placeholder':
                result = next(args_iter)
            elif node.op == 'get_attr':
                result = fetch_attr(node.target)
            elif node.op == 'call_function':
                result = node.target(*load_arg(node.args), **load_arg(node.kwargs))
            elif node.op == 'call_method':
                self_obj, *args = load_arg(node.args)
                kwargs = load_arg(node.kwargs)
                result = getattr(self_obj, node.target)(*args, **kwargs)
            elif node.op == 'call_module':
                result = self.modules[node.target](*load_arg(node.args), **load_arg(node.kwargs))

            # This is the only code specific to shape propagation.
            # you can delete this `if` branch and this becomes
            # a generic GraphModule interpreter.
            if isinstance(result, torch.Tensor):
                node.shape = result.shape
                node.dtype = result.dtype

            env[node.name] = result

        return load_arg(self.graph.result)

Как видите, полноценный интерпретатор для FX несложен, но может быть очень полезным. Для упрощения использования этого шаблона мы предоставляем класс Interpreter, который обобщает вышеупомянутую логику таким образом, что определённые аспекты выполнения интерпретатора могут быть переопределены посредством переопределения методов.

Помимо выполнения операций, мы также можем сгенерировать новый Graph, передавая значения Proxy через интерпретатор. Аналогично, мы предоставляем класс Transformer, который обобщает этот шаблон. Transformer ведет себя аналогично Interpreter, но вместо вызова метода run для получения конкретного выходного значения из модуля, вы вызываете метод Transformer.transform() для возврата нового GraphModule, который был подвергнут любым правилам трансформации, которые вы установили как переопределенные методы.

Примеры шаблона интерпретатора

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

Отладка

Введение

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

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

Распространённые ошибки при создании преобразований

  • Недетерминированный порядок итерирования set. В Python тип данных set неупорядочен. Использование set для хранения коллекций объектов, таких как Node , например, может привести к неожиданной недетерминированности. Пример — итерация по набору Node для их вставки в 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
"""

Здесь мы попытались проверить равенство значений двух глубоких нейронных моделей с оператором == равенства. Однако это некорректно, так как оператор возвращает тензор, а не булево значение, и для сравнения значений с плавающей точкой необходимо использовать погрешность (или ε), чтобы учесть некоммутативность операций с плавающей точкой (подробнее об этом см. здесь). Вместо этого мы можем использовать torch.allclose(), который даст приблизительное сравнение с учётом относительной и абсолютной пороговой погрешности:

assert torch.allclose(resnet18(input_image), transformed_resnet18(input_image))

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

Отладка сгенерированного кода

Так как FX генерирует функцию forward() для модулей GraphModule, использование традиционных методов отладки, таких как операторы print или pdb, не так просто. К счастью, у нас есть несколько методов отладки сгенерированного кода.

Использование pdb

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

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

def my_pass(inp: torch.nn.Module, tracer_class : type = fx.Tracer) -> torch.nn.Module:
    graph = tracer_class().trace(inp)
    # Transformation logic here
    # <...>

    # Return new Module
    return fx.GraphModule(inp, graph)

my_module = models.resnet18()
my_module_transformed = my_pass(my_module)

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

# When this line is executed at runtime, we will be dropped into an
# interactive `pdb` prompt. We can use the `step` or `s` command to
# step into the execution of the next line
import pdb; pdb.set_trace()

my_module_transformed(input_value)

Вывести сгенерированный код

Если вам нужно запускать один и тот же код несколько раз, то проход к нужному коду с помощью pdb может быть довольно утомительным. В этом случае один из подходов — просто скопировать и вставить сгенерированный forward проход в свой код и изучить его оттуда.

# Assume that `traced` is a GraphModule that has undergone some
# number of transforms

# Copy this code for later
print(traced)
# Print the code generated from symbolic tracing. This outputs:
"""
def forward(self, y):
    x = self.x
    add_1 = x + y;  x = y = None
    return add_1
"""

# Subclass the original Module
class SubclassM(M):
    def __init__(self):
        super().__init__()

    # Paste the generated `forward` function (the one we printed and
    # copied above) here
    def forward(self, y):
        x = self.x
        add_1 = x + y;  x = y = None
        return add_1

# Create an instance of the original, untraced Module. Then, create an
# instance of the Module with the copied `forward` function. We can
# now compare the output of both the original and the traced version.
pre_trace = M()
post_trace = SubclassM()

Использование функции to_folder из модуля GraphModule

GraphModule.to_folder() — это метод в GraphModule, который позволяет выгрузить сгенерированный код FX в папку. Хотя копирование прохода вперёд в код часто достаточно, как в Вывести сгенерированный код, может быть проще изучить модули и параметры с помощью to_folder.

m = symbolic_trace(M())
m.to_folder("foo", "Bar")
from foo import Bar
y = Bar()

После запуска приведенного выше примера, мы можем посмотреть на код внутри foo/module.py и изменить его по своему желанию (например, добавив print инструкции или используя pdb) для отладки сгенерированного кода.

Отладка преобразования

Теперь, когда мы определили, что преобразование создаёт неверный код, пришло время отладить само преобразование. Сначала мы проверим раздел Ограничения символического трассирования в документации. После того, как мы убедимся, что трассирование работает как ожидается, цель заключается в том, чтобы выяснить, что пошло не так во время нашего GraphModule преобразования. Быстрый ответ может быть в разделе Написание преобразований, но, если нет, существуют несколько способов изучить наш отслеженный модуль:

# Sample Module
class M(torch.nn.Module):
    def forward(self, x, y):
        return x + y

# Create an instance of `M`
m = M()

# Symbolically trace an instance of `M` (returns a GraphModule). In
# this example, we'll only be discussing how to inspect a
# GraphModule, so we aren't showing any sample transforms for the
# sake of brevity.
traced = symbolic_trace(m)

# Print the code produced by tracing the module.
print(traced)
# The generated `forward` function is:
"""
def forward(self, x, y):
    add = x + y;  x = y = None
    return add
"""

# Print the internal Graph.
print(traced.graph)
# This print-out returns:
"""
graph():
    %x : [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 команды отладчика для пошагового перемещения по вашей запущенной программе. Обычно вы устанавливаете точку останова (b LINE-NUMBER) при запуске pdb, а затем вызываете c для запуска программы до этой точки. Это позволит вам избежать прохождения по каждой строке выполнения (с помощью s или n) для достижения желаемой части кода. В качестве альтернативы, вы можете написать import pdb; pdb.set_trace() перед строкой, на которой хотите установить точку останова. Если вы добавите pdb.set_trace(), ваша программа автоматически запустится в режиме отладки при запуске. (Другими словами, вы можете просто ввести python FILENAME.py в командной строке вместо python -m pdb FILENAME.py.) После запуска файла в режиме отладки вы можете переходить по коду и изучать внутреннее состояние вашей программы с помощью определённых команд. В интернете есть множество отличных учебников по pdb , включая учебник RealPython «Python Debugging With Pdb».

IDE, такие как PyCharm или VSCode, обычно имеют встроенный отладчик. В вашей IDE вы можете либо а) использовать pdb, открыв окно терминала в вашей IDE (например, View → Terminal в VSCode), либо б) использовать встроенный отладчик (обычно графический оболочка вокруг pdb).

Ограничения символического трассирования

FX использует систему символического трассирования (также известного как символическое выполнение) для захвата семантики программ в преобразуемой/анализируемой форме. Система является трассированием, так как она выполняет программу (на самом деле torch.nn.Module или функцию) для записи операций. Она является символической, так как данные, протекающие через программу во время этого выполнения, не являются реальными данными, а символами (Proxy в терминологии FX).

Хотя символическое трассирование работает для большинства кода нейронных сетей, оно имеет некоторые ограничения.

Динамическое управление потоком

Основное ограничение символического трассирования заключается в том, что оно в настоящее время не поддерживает динамическое управление потоком. То есть циклы или if операторы, где условие может зависеть от входных значений программы.

Например, рассмотрим следующую программу:

def func_to_trace(x):
    if x.sum() > 0:
        return torch.relu(x)
    else:
        return torch.neg(x)

traced = torch.fx.symbolic_trace(func_to_trace)
"""
  <...>
  File "dyn.py", line 6, in func_to_trace
    if x.sum() > 0:
  File "pytorch/torch/fx/proxy.py", line 155, in __bool__
    return self.tracer.to_bool(self)
  File "pytorch/torch/fx/proxy.py", line 85, in to_bool
    raise TraceError('symbolically traced variables cannot be used as inputs to control flow')
torch.fx.proxy.TraceError: symbolically traced variables cannot be used as inputs to control flow
"""

Условие if оператора зависит от значения x.sum(), которое зависит от значения x, входной функции. Поскольку x может изменяться (то есть, если вы передаёте новый тензор ввода в отслеживаемую функцию), это динамическое управление потоком. Трассировка возвращается вверх по вашему коду, чтобы показать, где возникает эта ситуация.

Статическое управление потоком

С другой стороны, поддерживается так называемое статическое управление потоком. Статический поток управления — это циклы или if операторы, значение которых не может изменяться при вызовах. Как правило, в программах PyTorch такой поток управления возникает для кода, принимающего решения об архитектуре модели на основе гиперпараметров. В качестве конкретного примера:

import torch
import torch.fx

class MyModule(torch.nn.Module):
    def __init__(self, do_activation : bool = False):
        super().__init__()
        self.do_activation = do_activation
        self.linear = torch.nn.Linear(512, 512)

    def forward(self, x):
        x = self.linear(x)
        # This if-statement is so-called static control flow.
        # Its condition does not depend on any input values
        if self.do_activation:
            x = torch.relu(x)
        return x

without_activation = MyModule(do_activation=False)
with_activation = MyModule(do_activation=True)

traced_without_activation = torch.fx.symbolic_trace(without_activation)
print(traced_without_activation.code)
"""
def forward(self, x):
    linear_1 = self.linear(x);  x = None
    return linear_1
"""

traced_with_activation = torch.fx.symbolic_trace(with_activation)
print(traced_with_activation.code)
"""
import torch
def forward(self, x):
    linear_1 = self.linear(x);  x = None
    relu_1 = torch.relu(linear_1);  linear_1 = None
    return relu_1
"""

Оператор if if self.do_activation не зависит от входных данных функции, следовательно, он статический. do_activation можно рассматривать как гиперпараметр, и трассировки различных экземпляров Module с разными значениями этого параметра имеют разный код. Это допустимый шаблон, поддерживаемый символическим трассированием.

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

    • При использовании функционалов, таких как 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) [source]

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]) – Модуль или функция, которые нужно отследить и преобразовать в представление графа.
  • concrete_args (Optional[Dict[str, any]]) – Входные данные для частичной специализации
Возвращает

Модуль, созданный из записанных операций из root.

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

GraphModule

Примечание

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

torch.fx.wrap(fn_or_name) [source]

Эта функция может вызываться на уровне модуля, чтобы зарегистрировать fn_or_name как «листовую функцию». «Листовая функция» будет сохранена как узел CallFunction в трассировке FX вместо того, чтобы быть прослеженной:

# foo/bar/baz.py
def my_custom_function(x, y):
    return x * x + y * y

torch.fx.wrap('my_custom_function')

def fn_to_be_traced(x, y):
    # When symbolic tracing, the below call to my_custom_function will be inserted into
    # the graph rather than tracing it.
    return my_custom_function(x, y)

Эта функция также может быть эквивалентно использована как декоратор:

# foo/bar/baz.py
@torch.fx.wrap
def my_custom_function(x, y):
    return x * x + y * y

Оборачиваемая функция может рассматриваться как «листовая функция», аналогично понятию «листовых модулей», то есть это функции, которые остаются вызовами в трассировке FX, а не прослеживаются.

Параметры

fn_or_name (Union[str, Callable]) – Функция или имя глобальной функции, которую нужно вставить в граф при ее вызове

Примечание

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

class torch.fx.GraphModule(*args, **kwargs) [source]

GraphModule — это модуль nn.Module, сгенерированный из fx.Graph. У GraphModule есть атрибут graph, а также атрибуты code и forward, сгенерированные из этого graph.

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

При повторной присвоении graph, code и forward будут автоматически перегенерированы. Однако, если вы редактируете содержимое graph без повторного присвоения самого атрибута graph, необходимо вызвать recompile() для обновления сгенерированного кода.

Примечание

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

__init__(root, graph, class_name='GraphModule') [source]

Создаёт GraphModule.

Параметры
  • root (Union[torch.nn.Module, Dict[str, Any]) – root может быть экземпляром nn.Module или словарем, сопоставляющим строки с любым типом атрибута. В случае, если root является модулем, любые ссылки на объекты, основанные на модулях (через полное имя), в поле target узлов графа будут скопированы из соответствующего места в иерархии модулей root в иерархию модулей GraphModule. В случае, если root является словарем, полное имя, найденное в поле target узла, будет найдено непосредственно в ключах словаря. Объект, сопоставленный в словаре, будет скопирован в соответствующее место в иерархии модулей GraphModule.
  • graph (Graph) – graph содержит узлы, которые должен использовать GraphModule для генерации кода
  • class_name (str) – name обозначает имя этого GraphModule для отладки. Если оно не задано, все сообщения об ошибках будут сообщать об их происхождении из GraphModule. Может быть полезно установить это на исходное имя root или имя, которое имеет смысл в контексте вашего преобразования.

Примечание

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

add_submodule(target, m) [source]

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

Это устанавливает пустые модули там, где их ещё нет, если они являются подпутями target.

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

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

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

bool

Примечание

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

property code: str

Возвращает Python-код, сгенерированный из Graph лежащего в основе этого GraphModule.

delete_all_unused_submodules() [source]

Удаляет все неиспользуемые подмодули из self.

Модуль считается «используемым», если истинно хотя бы одно из следующих условий: 1. У него есть дочерние модули, которые используются 2. Его forward вызывается напрямую через узел call_module 3. У него есть атрибут, не являющийся модулем, который используется узлом get_attr

Этот метод можно вызвать для очистки nn.Module без ручного вызова delete_submodule для каждого неиспользуемого подмодуля.

Примечание

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

delete_submodule(target) [source]

Удаляет указанный подмодуль из self.

Модуль не будет удалён, если target не является допустимой целью.

Параметры

target (str) – Полное квалифицированное строковое имя нового подмодуля (см. пример в nn.Module.get_submodule для указания полного квалифицированного имени).

Возвращает
Будет ли удалён подмодуль, на который ссылается заданная строка.

Возвращаемое значение False означает, что target не является допустимой ссылкой на подмодуль.

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

bool

Примечание

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

property graph: Graph

Возвращает Graph лежащий в основе этого GraphModule.

print_readable(print_output=True) [source]

Возвращает Python-код, сгенерированный для текущего GraphModule и его дочерних GraphModule.

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

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

recompile() [source]

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

Примечание

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

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

PythonCode

to_folder(folder, module_name='FxModule') [source]
Dumps out module to folder with module_name so that it can be

импортированный с from <folder> import <module_name>

Аргументы:

folder (Union[str, os.PathLike]): Папка для записи кода

module_name (str): Top-level name to use for the Module while

запись кода

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

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

class torch.fx.Graph(owning_module=None, tracer_cls=None, tracer_extras=None) [source]

Graph является основной структурой данных, используемой в промежуточном представлении FX. Она состоит из последовательности Node , каждая из которых представляет вызов (или другие синтаксические конструкции). Список Node , вместе взятый, образует допустимую функцию Python.

Например, следующий код

import torch
import torch.fx

class MyModule(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.param = torch.nn.Parameter(torch.rand(3, 4))
        self.linear = torch.nn.Linear(4, 5)

    def forward(self, x):
        return torch.topk(torch.sum(self.linear(x + self.linear.weight).relu(), dim=-1), 3)

m = MyModule()
gm = torch.fx.symbolic_trace(m)

Сгенерирует следующий граф:

print(gm.graph)
graph(x):
    %linear_weight : [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) [source]

Создает пустой граф.

Примечание

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

call_function(the_function, args=None, kwargs=None, type_expr=None) [source]

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

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

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

Узел

Примечание

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

Примечание

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

call_method(method_name, args=None, kwargs=None, type_expr=None) [source]

Вставляет 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.

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

Узел

Примечание

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

Примечание

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

call_module(module_name, args=None, kwargs=None, type_expr=None) [source]

Вставляет call_module Node в Graph. Узел call_module представляет вызов функции forward() для Module в иерархии Module.

Параметры
  • module_name (str) – Полное имя Module в иерархии Module для вызова. Например, если отслеживаемый Module имеет подмодуль с именем foo, который имеет подмодуль с именем bar, полное имя foo.bar должно быть передано как module_name для вызова этого модуля.
  • args (Optional[Tuple[Argument, ...]]) – Позиционные аргументы, которые должны быть переданы вызываемому методу. Обратите внимание, что это *не* должно включать аргумент self.
  • kwargs (Optional[Dict[str, Argument]]) – Именованные аргументы, которые должны быть переданы вызываемому методу
  • type_expr (Optional[Any]) – необязательная аннотация типа, представляющая Python-тип, который будет иметь результат этого узла.
Возвращает

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

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

Узел

Примечание

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

Примечание

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

create_node(op, target, args=None, kwargs=None, name=None, type_expr=None) [source]

Создает Node и добавляет его в Graph в текущей точке вставки. Обратите внимание, что текущую точку вставки можно установить с помощью Graph.inserting_before() и Graph.inserting_after().

Параметры
  • op (str) – код операции для этого узла. Один из 'call_function', 'call_method', 'get_attr', 'call_module', 'placeholder' или 'output'. Семантика этих кодов операций описана в Graph строке документации.
  • args (Optional[Tuple[Argument, ...]]) – кортеж аргументов этого узла.
  • kwargs (Optional[Dict[str, Argument]]) – именованные аргументы этого узла
  • name (Optional[str]) – необязательное строковое имя для Node. Это повлияет на имя значения, присвоенного в сгенерированном коде Python.
  • type_expr (Optional[Any]) – необязательная аннотация типа, представляющая Python-тип, который будет иметь результат этого узла.
Возвращает

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

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

Узел

Примечание

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

eliminate_dead_code() [source]

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

Возвращает

Указывает, был ли граф изменён в результате обработки.

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

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) [source]

Удаляет узел Node из Graph. Выбрасывает исключение, если в Graph всё ещё есть пользователи этого узла.

Параметры

to_erase (Node) – Узел Node для удаления из Graph.

Примечание

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

get_attr(qualified_name, type_expr=None) [source]

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

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

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

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

Node

Примечание

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

Примечание

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

graph_copy(g, val_map, return_output_node=False) [source]

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

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

Значение в self , которое теперь эквивалентно выходному значению в g, если g имел узел output . None в противном случае.

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

Optional[Union[Tuple[Any, …], List[Any], Dict[str, Any], slice, range, Node, str, int, float, bool, complex, dtype, Tensor, device, memory_format, layout, OpOverload]]

Примечание

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

inserting_after(n=None) [source]
Устанавливает точку, в которой методы create_node и связанные с ним методы будут вставлять в граф.

При использовании в операторе with это временно установит точку вставки, а затем восстановит её по выходе из оператора with:

with g.inserting_after(n):
    ... # inserting after node n
... # insert point restored to what it was previously
g.inserting_after(n) #  set the insert point permanently

Аргументы:

n (Optional[Node]): Узел, перед которым нужно вставить. Если None, то вставка произойдёт после

начала всего графа.

Возвращает:

Управляющий ресурс, который восстановит точку вставки при выходе из __exit__.

Примечание

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

inserting_before(n=None) [source]
Устанавливает точку, в которой методы create_node и связанные с ним методы будут вставлять в граф.

При использовании в операторе with это временно установит точку вставки, а затем восстановит её по выходе из оператора with:

with g.inserting_before(n):
    ... # inserting before node n
... # insert point restored to what it was previously
g.inserting_before(n) #  set the insert point permanently

Аргументы:

n (Optional[Node]): Узел, перед которым нужно вставить. Если None, то вставка произойдёт перед

началом всего графа.

Возвращает:

Управляющий ресурс, который восстановит точку вставки при выходе из __exit__.

Примечание

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

lint() [source]

Выполняет различные проверки этого графа, чтобы убедиться, что он правильно сформирован. В частности: - Проверяет, что узлы имеют правильную принадлежность (принадлежат этому графу) - Проверяет, что узлы появляются в топологическом порядке - Если у этого графа есть владелец GraphModule, проверяет, что цели существуют в этом GraphModule

Примечание

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

node_copy(node, arg_transform=<function Graph.<lambda>>) [source]

Копирование узла из одной графы в другую. arg_transform необходимо преобразовать аргументы из графы узла в графу self. Пример:

# Copying all the nodes in `g` into `new_graph`
g : torch.fx.Graph = ...
new_graph = torch.fx.graph()
value_remap = {}
for node in g.nodes:
    value_remap[node] = new_graph.node_copy(node, lambda n : value_remap[n])
Параметры
  • node (Узел) – Узел для копирования в self.
  • arg_transform (Callable[[Узел], Argument]) – Функция, преобразующая Node аргументы в узле args и kwargs в эквивалентные аргументы в self. В простом случае, она должна извлекать значение из таблицы, сопоставляющей узлы исходной графы с self.
Тип возвращаемого значения

Узел

Примечание

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

property nodes: _node_list

Получение списка узлов, составляющих эту графу.

Обратите внимание, что это Node представление списка — это двусвязный список. Изменения во время итерации (например, удаление узла, добавление узла) безопасны.

Возвращаемое значение

Двусвязный список узлов. Обратите внимание, что reversed может быть вызван для изменения порядка итерации.

on_generate_code(make_transformer) [source]

Регистрация функции преобразования при генерации Python-кода

Аргументы:
make_transformer (Callable[[Optional[TransformCodeFunc]], TransformCodeFunc]):

Функция, возвращающая преобразователь кода для регистрации. Эта функция вызывается on_generate_code для получения преобразователя кода.

Эта функция также получает в качестве входных данных текущий зарегистрированный преобразователь кода (или None, если ничего не зарегистрировано), в случае, если не нужно его перезаписывать. Это полезно для объединения преобразователей кода.

Возвращаемое значение:

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

Пример:

gm: fx.GraphModule = ...

# This is a code transformer we want to register. This code
# transformer prepends a pdb import and trace statement at the very
# beginning of the generated torch.fx code to allow for manual
# debugging with the PDB library.
def insert_pdb(body):
    return ["import pdb; pdb.set_trace()\n", *body]

# Registers `insert_pdb`, and overwrites the current registered
# code transformer (given by `_` to the lambda):
gm.graph.on_generate_code(
    lambda _: insert_pdb
)

# Or alternatively, registers a code transformer which first
# runs `body` through existing registered transformer, then
# through `insert_pdb`:
gm.graph.on_generate_code(
    lambda current_trans: (
        lambda body: insert_pdb(
            current_trans(body) if current_trans
            else body
        )
    )
)

gm.recompile()
gm(*inputs)  # drops into pdb

Эта функция также может использоваться в качестве менеджера контекста с преимуществом автоматического восстановления ранее зарегистрированного преобразователя кода:

# ... continue from previous example

with gm.graph.on_generate_code(lambda _: insert_pdb):
    # do more stuff with `gm`...
    gm.recompile()
    gm(*inputs)  # drops into pdb

# now previous code transformer is restored (but `gm`'s code with pdb
# remains - that means you can run `gm` with pdb here too, until you
# run next `recompile()`).

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

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

output(result, type_expr=None) [source]

Вставка output Node в Graph. Узел output представляет оператор return в Python-коде. result — это значение, которое должно быть возвращено.

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

Примечание

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

Примечание

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

placeholder(name, type_expr=None, default_value) [source]

Вставка узла placeholder в графу. Узел placeholder представляет вход функции.

Параметры
  • name (str) – Имя входного значения. Соответствует имени позиционного аргумента функции, которую представляет этот Graph.
  • type_expr (Optional[Any]) – необязательная аннотация типа, представляющая Python-тип, который будет у результата этого узла. Это необходимо в некоторых случаях для правильной генерации кода (например, при последующем использовании функции в TorchScript).
  • default_value (Any) – Значение по умолчанию для этого аргумента функции. ПРИМЕЧАНИЕ: чтобы None можно было использовать в качестве значения по умолчанию, inspect.Signature.empty должно быть передано в качестве этого аргумента, чтобы указать, что параметр _не_ имеет значения по умолчанию.
Тип возвращаемого значения

Узел

Примечание

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

Примечание

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

print_tabular() [source]

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

Примечание

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

process_inputs(*args) [source]

Обработка аргументов для их передачи в граф FX.

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

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

process_outputs(out) [source]

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

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

python_code(root_module, *, verbose=False) [source]

Преобразование этой Graph в действительный Python-код.

Параметры

root_module (str) – Имя корневого модуля, в котором будут искать целевые имена.

Возвращаемое значение

src: исходный Python-код, представляющий объект globals: словарь глобальных имен в src -> соответствующие им объекты.

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

Объект PythonCode, состоящий из двух полей

Примечание

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

set_codegen(codegen) [source]

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

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

class torch.fx.Node(graph, name, op, target, args, kwargs, return_type=None) [source]

Node — это структура данных, представляющая отдельные операции в Graph. В основном, узлы представляют вызовы различных сущностей, таких как операторы, методы и модули (некоторые исключения включают узлы, определяющие входные и выходные данные функции). У каждого Node указана функция, определённая свойством op. Семантика Node для каждого значения op следующие:

  • placeholder представляет вход функции. Атрибут name определяет имя, которое получит это значение. target аналогичным образом — имя аргумента. args содержит либо: 1) ничего, либо 2) один аргумент, обозначающий параметр по умолчанию для входных данных функции. kwargs — это игнорируемое значение. Заполнители соответствуют параметрам функции (например, x) в выводе графа.
  • get_attr извлекает параметр из иерархии модулей. name аналогичным образом — имя, присваиваемое результату извлечения. target — полностью квалифицированное имя позиции параметра в иерархии модулей. args и kwargs — это игнорируемые значения.
  • call_function применяет свободную функцию к некоторым значениям. name аналогичным образом — имя присваиваемого значения. target — это функция, которая должна быть применена. args и kwargs представляют аргументы функции, следуя соглашениям Python для вызовов функций.
  • call_module применяет метод forward() модуля в иерархии модулей к заданным аргументам. name — как и раньше. target — это полностью квалифицированное имя модуля в иерархии модулей для вызова. args и kwargs представляют аргументы для вызова модуля, *исключая аргумент self*.
  • call_method вызывает метод на значении. name — как и раньше. target — строковое имя метода, который нужно применить к аргументу self. args и kwargs представляют аргументы для вызова модуля, *включая аргумент self*.
  • output содержит результат прослеживаемой функции в атрибуте args[0]. Это соответствует инструкции «return» в выводе графа.

Примечание

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

property all_input_nodes: List[Node]

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

Возвращает

Список Nodes , которые появляются в args и kwargs этого Node, в указанном порядке.

append(x) [source]

Вставляет x после этого узла в список узлов в графе. Эквивалентно self.next.prepend(x)

Параметры

x (Узел) – Узел, который следует поместить после этого узла. Должен быть членом того же графа.

Примечание

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

property args: Tuple[Optional[Union[Tuple[Any, ...], List[Any], Dict[str, Any], slice, range, Node, str, int, float, bool, complex, dtype, Tensor, device, memory_format, layout, OpOverload]], ...]

Кортеж аргументов для этого Node. Интерпретация аргументов зависит от кода операции узла. Дополнительную информацию см. в документации Node.

Присваивание этому свойству разрешено. Все учёт использования и пользователей обновляется автоматически при присваивании.

format_node(placeholder_names=None, maybe_return_typename=None) [source]

Возвращает строковое описание узла self. Это метод может использоваться без аргументов в качестве инструмента отладки.

Этот метод также используется внутри метода __str__ Graph . Вместе строки в placeholder_names и maybe_return_typename составляют подпись автоматически генерируемой функции forward в окружающем GraphModule этого графа. placeholder_names и maybe_return_typename не должны использоваться в иных случаях.

Параметры
  • placeholder_names (Optional[Список[строка]]) – Список, который будет хранить отформатированные строки, представляющие заполнители в сгенерированной функции forward. Только внутреннее использование.
  • maybe_return_typename (Optional[Список[строка]]) – Одноэлементный список, который будет хранить отформатированную строку, представляющую результат сгенерированной функции forward . Только внутреннее использование.
Возвращает
If 1) we’re using format_node as an internal helper

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

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

строка

Примечание

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

is_impure() [source]

Возвращает, является ли эта операция нечистой, т.е. если её операция является заполнителем или выходом, или если вызов функции или вызов модуля нечистый.

Возвращает

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

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

булево

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

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

property kwargs: Dict[str, Optional[Union[Tuple[Any, ...], List[Any], Dict[str, Any], slice, range, Node, str, int, float, bool, complex, dtype, Tensor, device, memory_format, layout, OpOverload]]]

Словарь ключевых аргументов для этого Node. Интерпретация аргументов зависит от кода операции узла. Дополнительную информацию см. в документации Node.

Присваивание этому свойству разрешено. Все учёт использования и пользователей обновляется автоматически при присваивании.

property next: Node

Возвращает следующий Node в связанном списке узлов.

Возвращает

Следующий Node в связанном списке узлов.

normalized_arguments(root, arg_types=None, kwarg_types=None, normalize_to_only_use_kwargs=False) [source]

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

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

Optional[ArgsKwargsPair]

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

Этот API находится в стадии разработки и НЕ совместим с предыдущими версиями.

prepend(x) [source]

Вставить x перед этим узлом в список узлов в графе. Пример:

Before: p -> self
        bx -> x -> ax
After:  p -> x -> self
        bx -> ax
Параметры

x (Узел) – Узел, который нужно поместить перед этим узлом. Должен быть членом той же схемы.

Примечание

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

property prev: Node

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

Возвращает

Предыдущий Node в связанном списке узлов.

replace_all_uses_with(replace_with, delete_user_cb=<function Node.<lambda>>, *, propagate_meta=False) [source]

Заменить все использования self в схеме на узел replace_with.

Параметры
  • replace_with (Узел) – Узел, на который нужно заменить все использования self.
  • delete_user_cb (Callable) – Обратный вызов, который вызывается для определения, следует ли удалять данного пользователя узла self.
  • propagate_meta (bool) – Нужно ли копировать все свойства в поле .meta исходного узла на узел-замену. Для безопасности это допустимо только в том случае, если у узла-замены нет существующего поля .meta.
Возвращает

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

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

Список[Узел]

Примечание

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

replace_input_with(old_input, new_input) [source]

Перебрать входные узлы self, и заменить все экземпляры old_input на new_input.

Параметры
  • old_input (Узел) – Узел старого входного значения, подлежащий замене.
  • new_input (Узел) – Новый входной узел, на который нужно заменить old_input.

Примечание

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

property stack_trace: Optional[str]

Возвращает трассировку стека Python, записанную во время трассировки, если таковая имеется. При трассировке с помощью fx.Tracer это свойство обычно заполняется Tracer.create_proxy. Чтобы записывать трассировки стека во время трассировки для отладки, установите record_stack_traces = True в экземпляре Tracer. При трассировке с помощью dynamo это свойство будет заполняться по умолчанию OutputGraph.create_proxy.

stack_trace будет содержать внутренний кадр в конце строки.

update_arg(idx, arg) [source]

Обновить существующий позиционный аргумент, чтобы он содержал новое значение arg. После вызова self.args[idx] == arg.

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

Примечание

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

update_kwarg(key, arg) [source]

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

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

Примечание

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

class torch.fx.Tracer(autowrap_modules=(math,), autowrap_functions=()) [source]

Tracer — это класс, реализующий функциональность символического трассирования в torch.fx.symbolic_trace. Вызов symbolic_trace(m) эквивалентен вызову Tracer().trace(m).

Класс Tracer можно наследоваться, чтобы переопределить различные аспекты процесса трассирования. Различные переопределяемые аспекты описаны в документации методов данного класса.

Примечание

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

call_module(m, forward, args, kwargs) [source]

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

По умолчанию, поведение заключается в проверке, является ли вызываемый модуль листовым модулем с помощью is_leaf_module. Если да, то генерируется узел call_module, ссылающийся на m в Graph. В противном случае, вызывается Module стандартным образом, прослеживая операции в его методе forward.

Этот метод можно переопределить, например, для создания вложенных отслеживаемых GraphModules или для любого другого поведения, необходимого при трассировании через границы Module.

Параметры
  • m (Модуль) — Модуль, для которого генерируется вызов.
  • forward (Callable) — Метод forward() вызываемого Module.
  • args (Tuple) — Аргументы вызова модуля.
  • kwargs (Dict) — Имена аргументов вызова модуля.
Возвращает

Значение, возвращаемое вызовом модуля. В случае генерации узла call_module, это значение типа Proxy . В противном случае, это значение, возвращённое вызовом Module.

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

Any

Примечание

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

create_arg(a) [source]

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

По умолчанию, поведение включает:

  1. Итерацию по коллекционным типам (например, кортеж, список, словарь) и рекурсивный вызов create_args для элементов.
  2. Для объекта Proxy возвращает ссылку на подлежащий IR Node.
  3. Для объекта тензора, не являющегося Proxy, генерирует IR для различных случаев:

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

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

Параметры

a (Any) — Значение, которое должно быть сгенерировано как Argument в Graph.

Возвращает

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

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

Optional[Union[Tuple[Any, …], List[Any], Dict[str, Any], slice, range, Node, str, int, float, bool, complex, dtype, Tensor, device, memory_format, layout, OpOverload]]

Примечание

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

create_args_for_root(root_fn, is_module, concrete_args=None) [source]

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

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

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

create_node(kind, target, args, kwargs, name=None, type_expr=None)

Вставляет узел графа, используя target, args, kwargs и имя.

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

Примечание

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

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

Node

create_proxy(kind, target, args, kwargs, name=None, type_expr=None, proxy_factory_fn=None)

Создаёт узел из заданных аргументов, затем возвращает узел, обернутый в объект Proxy.

Если kind = ‘placeholder’, то мы создаём узел, представляющий параметр функции. Если нам нужно закодировать параметр по умолчанию, мы используем кортеж args . args в противном случае пуст для узлов placeholder.

Примечание

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

getattr(attr, attr_val, parameter_proxy_cache) [source]

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

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

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

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

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

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

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

is_leaf_module(m, module_qualified_name) [source]

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

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

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

логическое значение

Примечание

Совместимость с предыдущими версиями для этого API гарантирована.

iter(obj)
Вызывается при итерировании по объекту-прокси, например,

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

Примечание

Совместимость с предыдущими версиями для этого API гарантирована.

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

Итератор

keys(obj)
Вызывается при вызове метода keys() у объекта-прокси.

Это происходит при вызове ** у прокси. Это должно вернуть итератор, если ** должен работать в вашем пользовательском трейсере.

Примечание

Совместимость с предыдущими версиями для этого API гарантирована.

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

любой

path_of_module(mod) [source]

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

Параметры

mod (строка) – Модуль для которого необходимо получить квалифицированное имя.

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

строка

Примечание

Совместимость с предыдущими версиями для этого API гарантирована.

proxy(node)

Примечание

Совместимость с предыдущими версиями для этого API гарантирована.

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

Прокси

to_bool(obj)
Вызывается при преобразовании объекта-прокси в булевое значение, например,

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

Примечание

Совместимость с предыдущими версиями для этого API гарантирована.

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

логическое значение

trace(root, concrete_args=None) [source]

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

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

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

FX-граф, представляющий семантику переданного root.

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

Граф

Примечание

Совместимость с предыдущими версиями для этого API гарантирована.

class torch.fx.Proxy(node, tracer=None) [source]

Объекты Proxy являются Node обёртками, которые проходят через программу во время символического отслеживания и записывают все операции (torch вызовы функций, вызовы методов, операторы), с которыми они взаимодействуют, в растущий FX-граф.

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

Объекты Proxy не могут быть итерированы. Другими словами, символический трейсер выдаст ошибку, если Proxy используется в цикле или в качестве аргумента функции *args/**kwargs.

Есть два основных способа обойти это: 1. Вынести неотслеживаемую логику в функцию верхнего уровня и применить fx.wrap к ней. 2. Если управляющая структура статична (т. е. количество циклов основано на некотором гиперпараметре), код можно оставить на своем месте и переработать в нечто вроде:

for i in range(self.some_hyperparameter):
    indexed_item = proxied_value[i]

Для более подробного описания внутренних механизмов прокси, обратитесь к разделу «Прокси» в torch/fx/OVERVIEW.md

Примечание

Совместимость с предыдущими версиями для этого API гарантирована.

class torch.fx.Interpreter(module, garbage_collect_values=True) [source]

Интерпретатор выполняет узел FX-графа по узлу. Этот шаблон может быть полезен для многих задач, включая написание преобразований кода и проход анализа.

Методы в классе Interpreter могут быть переопределены для настройки поведения выполнения. Карта переопределяемых методов с точки зрения иерархии вызовов:

run()
    +-- run_node
        +-- placeholder()
        +-- get_attr()
        +-- call_function()
        +-- call_method()
        +-- call_module()
        +-- output()

Пример

Предположим, что мы хотим поменять все экземпляры torch.neg на torch.sigmoid и наоборот (включая их эквиваленты метода Tensor). Мы можем создать подкласс Interpreter следующим образом:

class NegSigmSwapInterpreter(Interpreter):
    def call_function(self, target : Target,
                      args : Tuple, kwargs : Dict) -> Any:
        if target == torch.sigmoid:
            return torch.neg(*args, **kwargs)
        return super().call_function(n)

    def call_method(self, target : Target,
                    args : Tuple, kwargs : Dict) -> Any:
        if target == 'neg':
            call_self, *args_tail = args
            return call_self.sigmoid(*args_tail, **kwargs)
        return super().call_method(n)

def fn(x):
    return torch.sigmoid(x).neg()

gm = torch.fx.symbolic_trace(fn)
input = torch.randn(3, 4)
result = NegSigmSwapInterpreter(gm).run(input)
torch.testing.assert_close(result, torch.neg(input).sigmoid())
Параметры
  • module (GraphModule) – Модуль, подлежащий выполнению
  • garbage_collect_values (bool) – Удалять ли значения после их последнего использования в рамках выполнения модуля. Это гарантирует оптимальное использование памяти во время выполнения. Это можно отключить, чтобы, например, просмотреть все промежуточные значения в выполнении, посмотрев на атрибут Interpreter.env.

Примечание

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

boxed_run(args_list) [source]

Выполнить module посредством интерпретации и вернуть результат. Это использует соглашение вызова «boxed», где вы передаёте список аргументов, который будет очищен интерпретатором. Это гарантирует, что входные тензоры будут немедленно освобождены.

Примечание

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

call_function(target, args, kwargs) [source]

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

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

Any

Возвращаемое значение

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

Примечание

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

call_method(target, args, kwargs) [source]

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

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

Any

Возвращаемое значение

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

Примечание

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

call_module(target, args, kwargs) [source]

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

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

Any

Возвращаемое значение

Any: значение, возвращённое вызовом модуля

Примечание

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

fetch_args_kwargs_from_env(n) [source]

Извлечь конкретные значения args и kwargs узла n из текущей среды выполнения.

Параметры

n (Node) – Узел, для которого необходимо извлечь args и kwargs.

Возвращаемое значение

args и kwargs с конкретными значениями для n.

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

Tuple[Tuple, Dict]

Примечание

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

fetch_attr(target) [source]

Извлечь атрибут из иерархии Module self.module.

Параметры

target (str) – Полное квалифицированное имя атрибута для извлечения

Возвращаемое значение

Значение атрибута.

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

Any

Примечание

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

get_attr(target, args, kwargs) [source]

Выполнить узел get_attr. Извлечёт значение атрибута из иерархии Module self.module.

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

Значение извлечённого атрибута.

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

Any

Примечание

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

map_nodes_to_values(args, n) [source]

Рекурсивно проходит по args и ищет конкретное значение для каждого Node в текущей среде выполнения.

Параметры
  • args (Argument) – Структура данных для поиска конкретных значений
  • n (Узел) – Узел, к которому args относится. Используется только для отладки ошибок.
Тип возвращаемого значения

Optional[Union[Кортеж[Любой, …], Список[Любой], Словарь[строка, Любой], срез, диапазон, Узел, строка, целое число, вещественное число, логическое значение, комплексное число, тип данных, тензор, устройство, форма памяти, структура, OpOverload]]

Примечание

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

output(target, args, kwargs) [source]

Выполняет узел output. Это просто извлекает значение, на которое ссылается узел output, и возвращает его.

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

Возвращаемое значение, на которое указывает выходной узел

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

Любой

Примечание

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

placeholder(target, args, kwargs) [source]

Выполняет узел placeholder. Обратите внимание, что это состояние: Interpreter поддерживает внутренний итератор по аргументам, переданным в run, и этот метод возвращает next() для этого итератора.

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

Значение аргумента, которое было получено.

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

Любой

Примечание

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

run(*args, initial_env=None, enable_io_processing=True) [source]

Выполняет module через интерпретацию и возвращает результат.

Параметры
  • *args – Аргументы модуля для выполнения в позиционном порядке
  • initial_env (Optional[Dict[Узел, Любой]]) – Необязательная начальная среда выполнения. Это словарь, сопоставляющий Node любому значению. Это можно использовать, например, для предварительной заполнения результатов для определенных Nodes для частичной оценки в интерпретаторе.
  • enable_io_processing (bool) – Если True, мы сначала обрабатываем входные и выходные данные функциями process_inputs и process_outputs графа, прежде чем использовать их.
Возвращает

Значение, возвращенное при выполнении модуля

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

Любой

Примечание

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

run_node(n) [source]

Выполняет конкретный узел n и возвращает результат. Вызывает placeholder, get_attr, call_function, call_method, call_module или output в зависимости от node.op

Параметры

n (Узел) – Узел для выполнения

Возвращает

Результат выполнения n

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

Любой

Примечание

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

class torch.fx.Transformer(module) [source]

Transformer — это специальный тип интерпретатора, который создаёт новый Module. Он предоставляет метод transform(), который возвращает преобразованный Module. Transformer не требует аргументов для выполнения, в отличие от Interpreter. Transformer работает полностью символически.

Пример

Предположим, мы хотим поменять все экземпляры torch.neg на torch.sigmoid и наоборот (включая эквиваленты методов Tensor). Мы можем сделать это, наследуя класс Transformer:

class NegSigmSwapXformer(Transformer):
    def call_function(self, target : 'Target', args : Tuple[Argument, ...], kwargs : Dict[str, Any]) -> Any:
        if target == torch.sigmoid:
            return torch.neg(*args, **kwargs)
        return super().call_function(n)

    def call_method(self, target : 'Target', args : Tuple[Argument, ...], kwargs : Dict[str, Any]) -> Any:
        if target == 'neg':
            call_self, *args_tail = args
            return call_self.sigmoid(*args_tail, **kwargs)
        return super().call_method(n)

def fn(x):
    return torch.sigmoid(x).neg()

gm = torch.fx.symbolic_trace(fn)

transformed : torch.nn.Module = NegSigmSwapXformer(gm).transform()
input = torch.randn(3, 4)
torch.testing.assert_close(transformed(input), torch.neg(input).sigmoid())
Параметры

module (GraphModule) — Преобразуемый Module.

Примечание

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

call_function(target, args, kwargs) [source]

Примечание

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

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

Any

call_module(target, args, kwargs) [source]

Примечание

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

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

Any

get_attr(target, args, kwargs) [source]

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

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

Proxy

Примечание

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

placeholder(target, args, kwargs) [source]

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

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

Proxy

Примечание

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

transform() [source]

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

Примечание

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

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

GraphModule

torch.fx.replace_pattern(gm, pattern, replacement) [source]

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

Параметры
  • gm (GraphModule) — GraphModule, который оборачивает Graph для обработки
  • pattern (Union[Callable, GraphModule]) — Подграф для поиска в gm для замены
  • replacement (Union[Callable, 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):
        super().__init__()

    def forward(self, x, w1, w2):
        m1 = torch.cat([w1, w2]).sum()
        m2 = torch.cat([w1, w2]).sum()
        return x + torch.max(m1) + torch.max(m2)

def pattern(w1, w2):
    return torch.cat([w1, w2]).sum()

def replacement(w1, w2):
    return torch.stack([w1, w2])

traced_module = symbolic_trace(M())

subgraph_rewriter.replace_pattern(traced_module, pattern, replacement)

Этот код сначала найдет совпадение с pattern в методе forward класса traced_module. Поиск совпадений производится по отношениям использования и определения, а не по именам узлов. Например, если у вас есть p = torch.cat([a, b]) в pattern, вы можете найти совпадение с m = torch.cat([a, b]) в исходной функции forward, несмотря на то, что имена переменных разные (p vs m).

Выражение return в pattern ищется только по своему значению; оно может или не может совпасть с выражением return в более крупном графе. Другими словами, шаблон не должен распространяться до конца более крупного графа.

Когда шаблон найден, он удаляется из более крупной функции и заменяется replacement. Если в более крупной функции есть несколько совпадений с pattern, каждое непересекающееся совпадение будет заменено. В случае перекрытия совпадений, будет заменено первое найденное совпадение в наборе перекрывающихся совпадений. ("Первое" здесь определяется как первое в топологическом порядке отношений использования и определения узлов. В большинстве случаев, первый узел — это параметр, который появляется непосредственно после self, а последний узел — то, что функция возвращает.)

Важно отметить, что параметры pattern Callable должны использоваться в самом Callable, а параметры replacement Callable должны совпадать с шаблоном. Первый принцип объясняет, почему в приведенном коде функция forward имеет параметры x, w1, w2, но функция pattern только w1, w2. pattern не использует x, поэтому не должна указывать x как параметр. В качестве примера второго принципа рассмотрим замену

def pattern(x, y):
    return torch.neg(x) + torch.relu(y)

на

def replacement(x, y):
    return torch.relu(x)

В этом случае replacement нуждается в таком же количестве параметров, как и pattern (как x, так и y), даже если параметр y не используется в replacement.

После вызова subgraph_rewriter.replace_pattern сгенерированный Python-код выглядит так:

def forward(self, x, w1, w2):
    stack_1 = torch.stack([w1, w2])
    sum_1 = stack_1.sum()
    stack_2 = torch.stack([w1, w2])
    sum_2 = stack_2.sum()
    max_1 = torch.max(sum_1)
    add_1 = x + max_1
    max_2 = torch.max(sum_2)
    add_2 = add_1 + max_2
    return add_2

Примечание

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

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

Spec-Zone.ru

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