Spec-Zone.ru › PyTorch 1

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 : [#users=1] = placeholder[target=x]
    %param : [#users=1] = get_attr[target=param]
    %add : [#users=1] = call_function[target=operator.add](args = (%x, %param), kwargs = {})
    %linear : [#users=1] = call_module[target=linear](args = (%add,), kwargs = {})
    %clamp : [#users=1] = call_method[target=clamp](args = (%linear,), kwargs = {min: 0.0, max: 1.0})
    return clamp
"""

# Code generation - valid Python code
print(symbolic_traced.code)
"""
def forward(self, x):
    param = self.param
    add = x + param;  x = param = None
    linear = self.linear(add);  add = None
    clamp = linear.clamp(min = 0.0, max = 1.0);  linear = None
    return clamp
"""

Символический трассер выполняет «символическое выполнение» Python-кода. Он подаёт поддельные значения, называемые Proxy, через код. Операции с этими Proxy записываются. Дополнительную информацию о символическом прослеживании можно найти в документации symbolic_trace() и Tracer.

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

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

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

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

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

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

import torch
import torch.fx

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

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

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

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

Ваше преобразование примет на вход torch.nn.Module, получит Graph из него, внесет некоторые изменения и вернёт новый torch.nn.Module. Вы должны понимать, что 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

имя

цель

аргументы

ключевые аргументы

placeholder

x

x

()

{}

get_attr

linear_weight

linear.weight

()

{}

call_function

add_1

<встроенная функция add>

(x, linear_weight)

{}

call_module

linear_1

linear

(add_1,)

{}

call_method

relu_1

relu

(linear_1,)

{}

call_function

sum_1

<встроенный метод sum …>

(relu_1,)

{‘dim’: -1}

call_function

topk_1

<встроенный метод topk …>

(sum_1, 3)

{}

output

output

output

(topk_1,)

{}

Мы можем использовать эту информацию для ответа на поставленные выше вопросы.

  • Какие входные данные принимаются методом? В FX входные данные метода определяются с помощью специальных placeholder узлов. В данном случае у нас есть один placeholder узел с target значением x, что означает, что у нас есть один аргумент (без self) с именем x.
  • Какие операции выполняются внутри метода? get_attr, call_function, call_module, и call_method узлы представляют операции в методе. Подробное описание семантики всех этих узлов можно найти в документации Node.
  • Какое значение возвращается методом? Значение возвращаемое методом в Graph определяется специальным output узлом.

Теперь, зная основы представления кода в FX, мы можем изучить, как изменить Graph.

Изменение графа

Прямое изменение графа

Один из подходов к созданию нового Graph заключается в прямом изменении исходного. Для этого мы можем просто взять Graph, полученный из символьного прослеживания, и изменить его. Например, предположим, что нам нужно заменить вызовы torch.add() на вызовы torch.mul().

import torch
import torch.fx

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

def transform(m: torch.nn.Module,
              tracer_class : type = fx.Tracer) -> torch.nn.Module:
    graph : fx.Graph = tracer_class().trace(m)
    # FX represents its Graph as an ordered list of
    # nodes, so we can iterate through them.
    for node in graph.nodes:
        # Checks if we're calling a function (i.e:
        # torch.add)
        if node.op == 'call_function':
            # The target attribute is the function
            # that call_function calls.
            if node.target == torch.add:
                node.target = torch.mul

    graph.lint() # Does some checks to make sure the
                 # Graph is well-formed.

    return fx.GraphModule(m, graph)

Мы также можем выполнить более сложные переписывания Graph, такие как удаление или добавление узлов. Для облегчения этих преобразований FX предоставляет служебные функции для изменения графа, которые можно найти в документации Graph. Пример использования этих API для добавления вызова torch.relu() приведён ниже.

# Specifies the insertion point. Any nodes added to the
# Graph within this scope will be inserted after `node`
with traced.graph.inserting_after(node):
    # Insert a new `call_function` node calling `torch.relu`
    new_node = traced.graph.call_function(
        torch.relu, args=(node,))

    # We want all places that used the value of `node` to
    # now use that value after the `relu` call we've added.
    # We use the `replace_all_uses_with` API to do this.
    node.replace_all_uses_with(new_node)

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

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

FX также предоставляет ещё один уровень автоматизации поверх прямого изменения графа. API replace_pattern() по существу является инструментом «найти/заменить» для редактирования Graph. Он позволяет указать 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-арными (например, операция сложения является бинарной), так как вызов 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 : [#users=1] = placeholder[target=x]
    %y : [#users=1] = placeholder[target=y]
    %add : [#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» “Python Debugging With Pdb”.

IDE, такие как PyCharm или VSCode, обычно имеют встроенный отладчик. В вашей IDE вы можете выбрать либо a) использовать pdb, открыв окно терминала в вашей IDE (например, View → Terminal в VSCode), или b) использовать встроенный отладчик (обычно графический интерфейс вокруг 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 можно рассматривать как гиперпараметр, и трассировки разных экземпляров MyModule с различными значениями этого параметра имеют различный код. Это допустимый шаблон, который поддерживается символичным прослеживанием.

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

def f(x, flag):
    if flag: return x
    else: return x*2

fx.symbolic_trace(f) # Fails!

fx.symbolic_trace(f, concrete_args={'flag': True})

В случае истинного динамического управления потоком участки программы, содержащие этот код, можно прослеживать как вызовы метода (см. Настройка прослеживания с классом Tracer) или функции (см. wrap()) вместо прослеживания через них.

Не-torch функции

FX использует __torch_function__ в качестве механизма перехвата вызовов (см. технический обзор для получения дополнительной информации об этом). Некоторые функции, такие как встроенные функции Python или функции из модуля math, не покрываются __torch_function__, но мы все равно хотели бы захватить их в символическом прослеживании. Например:

import torch
import torch.fx
from math import sqrt

def normalize(x):
    """
    Normalize `x` by the size of the batch dimension
    """
    return x / sqrt(len(x))

# It's valid Python code
normalize(torch.rand(3, 4))

traced = torch.fx.symbolic_trace(normalize)
"""
  <...>
  File "sqrt.py", line 9, in normalize
    return x / sqrt(len(x))
  File "pytorch/torch/fx/proxy.py", line 161, in __len__
    raise RuntimeError("'len' is not supported in symbolic tracing by default. If you want "
RuntimeError: 'len' is not supported in symbolic tracing by default. If you want this call to be recorded, please call torch.fx.wrap('len') at module scope
"""

Ошибка сообщает нам, что встроенная функция len не поддерживается. Мы можем сделать так, чтобы функции такого рода записывались в трассировке как прямые вызовы, используя API wrap():

torch.fx.wrap('len')
torch.fx.wrap('sqrt')

traced = torch.fx.symbolic_trace(normalize)

print(traced.code)
"""
import math
def forward(self, x):
    len_1 = len(x)
    sqrt_1 = math.sqrt(len_1);  len_1 = None
    truediv = x / sqrt_1;  x = sqrt_1 = None
    return truediv
"""

Настройка прослеживания с классом Tracer

Класс Tracer является классом, лежащим в основе реализации symbolic_trace. Поведение прослеживания можно настроить, создав подкласс Tracer, например так:

class MyCustomTracer(torch.fx.Tracer):
    # Inside here you can override various methods
    # to customize tracing. See the `Tracer` API
    # reference
    pass


# Let's use this custom tracer to trace through this module
class MyModule(torch.nn.Module):
    def forward(self, x):
        return torch.relu(x) + torch.ones(3, 4)

mod = MyModule()

traced_graph = MyCustomTracer().trace(mod)
# trace() returns a Graph. Let's wrap it up in a
# GraphModule to make it runnable
traced = torch.fx.GraphModule(mod, traced_graph)

Листовые модули

Листовые модули — это модули, которые появляются в качестве вызовов в символьной трассировке, а не прослеживаются через них. Стандартный набор листовых модулей — это набор стандартных экземпляров модулей torch.nn. Например:

class MySpecialSubmodule(torch.nn.Module):
    def forward(self, x):
        return torch.neg(x)

class MyModule(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.linear = torch.nn.Linear(3, 4)
        self.submod = MySpecialSubmodule()

    def forward(self, x):
        return self.submod(self.linear(x))

traced = torch.fx.symbolic_trace(MyModule())
print(traced.code)
# `linear` is preserved as a call, yet `submod` is traced though.
# This is because the default set of "Leaf Modules" includes all
# standard `torch.nn` modules.
"""
import torch
def forward(self, x):
    linear_1 = self.linear(x);  x = None
    neg_1 = torch.neg(linear_1);  linear_1 = None
    return neg_1
"""

Набор листовых модулей можно настроить, переопределив Tracer.is_leaf_module().

Разное

  • Конструкторы тензоров (например, torch.zeros, torch.ones, torch.rand, torch.randn, torch.sparse_coo_tensor в настоящее время не прослеживаются.

    • Детерминированные конструкторы (zeros, ones) могут быть использованы, и значение, которое они производят, будет встроено в трассировку как константа. Это проблема только в том случае, если аргументы этих конструкторов ссылаются на динамические размеры ввода. В этом случае ones_like или zeros_like могут быть подходящей заменой.
    • Недетерминированные конструкторы (rand, randn будут содержать одно случайное значение, встроенное в трассировку. Вероятно, это не то поведение, которое ожидалось. Одним из решений является обертывание torch.randn в функцию torch.fx.wrap и вызов этой функции вместо этого.
    @torch.fx.wrap
    def torch_randn(x, shape):
        return torch.randn(shape)
    
    def f(x):
        return x + torch_randn(x, 5)
    fx.symbolic_trace(f)
    
    • Это поведение может быть исправлено в будущих выпусках.
  • Аннотации типов

    • Аннотации типов в стиле Python 3 (например, func(x : torch.Tensor, y : int) -> torch.Tensor) поддерживаются и будут сохранены символичным прослеживанием.
    • Аннотации типов в стиле комментариев Python 2 # type: (torch.Tensor, int) -> torch.Tensor в настоящее время не поддерживаются.
    • Аннотации локальных имен внутри функции в настоящее время не поддерживаются.
  • Особенности работы флага training и подмодулей

    • При использовании функционалов, таких как torch.nn.functional.dropout, часто аргумент обучения передается как 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_allclose(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(), флаг обучения инкапсулируется и — из-за сохранения модели объекта 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_allclose(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
Parameters:
  • root (Union[torch.nn.Module, Callable]) – Модуль или функция, которая будет прослеживаться и преобразовываться в представление графа.
  • concrete_args (Optional[Dict[str, any]]) – Входные данные для частичной специализации
Returns:

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

Return type:

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, а не прослеживаются через них.

Parameters:

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

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

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

Этот 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 : [#users=1] = self.linear.weight
    %add_1 : [#users=1] = call_function[target=operator.add](args = (%x, %linear_weight), kwargs = {})
    %linear_1 : [#users=1] = call_module[target=linear](args = (%add_1,), kwargs = {})
    %relu_1 : [#users=1] = call_method[target=relu](args = (%linear_1,), kwargs = {})
    %sum_1 : [#users=1] = call_function[target=torch.sum](args = (%relu_1,), kwargs = {dim: -1})
    %topk_1 : [#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. Узел вызова представляет вызов 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 узел.

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

Node

Примечание

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

Примечание

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

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

Вставляет узел вызова call_method Node в диаграмму Graph. Узел вызова метода представляет вызов заданного метода для 0-го элемента args.

Параметры:
  • method_name (str) – Имя метода, который нужно применить к аргументу self. Например, если args[0] — Node представляющий Tensor, то для вызова relu() на этом Tensor, передайте relu в method_name.
  • args (Optional[Tuple[Argument, ...]]) – Позиционные аргументы, которые должны быть переданы вызываемому методу. Обратите внимание, что это должно включать аргумент self.
  • kwargs (Optional[Dict[str, Argument]]) – Именованные аргументы, которые должны быть переданы вызываемому методу
  • type_expr (Optional[Any]) – необязательная аннотация типа, представляющая тип Python, который будет иметь результат этого узла.
Возвращает:

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

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

Node

Примечание

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

Примечание

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

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

Вставляет вызов call_module Node в диаграмму Graph. Узел вызова модуля представляет вызов функции 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 узел.

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

Node

Примечание

Для этого метода применяются те же правила вставки и выражения типов, что и для 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().

Parameters:
  • 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-тип, который будет иметь вывод этого узла.
Returns:

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

Return type:

Node

Примечание

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

eliminate_dead_code() [source]

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

Returns:

Изменился ли граф в результате прохода.

Return type:

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.

Parameters:

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

Примечание

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

get_attr(qualified_name, type_expr=None) [source]

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

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

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

Return type:

Node

Примечание

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

Примечание

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

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

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

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

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

Return type:

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

Примечание

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

property owning_module

Возвращает модуль, владеющий этим GraphModule, если такой существует. None если модуль-владелец отсутствует или если модулей-владельцев несколько.

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

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

Параметры:
  • name (строка) – Имя входного значения. Соответствует имени позиционного аргумента функции, которую представляет данный 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) – Имя корневого модуля, в котором нужно искать квалифицированные имена целевых объектов. Обычно это ‘self’.

Возвращает:

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, Node, str, int, float, bool, complex, dtype, Tensor, device, memory_format, layout]], ...]

Кортеж аргументов этого 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[Список[str]]) – Список, который будет хранить отформатированные строки, представляющие заполнители в сгенерированной функции forward. Использование только внутри.
  • maybe_return_typename (Optional[Список[str]]) – Список с одним элементом, который будет хранить отформатированную строку, представляющую выходную функцию forward. Использование только внутри.
Возвращает:
If 1) we’re using format_node as an internal helper

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

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

str

Примечание

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

is_impure() [source]

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

Возвращает:

Если операция нечистая или нет.

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

bool

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

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

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

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

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

property next: Node

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

Возвращает:

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

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

Возвращает нормализованные аргументы для Python-мишеней. Это означает, что args/kwargs будут сопоставлены с сигнатурой модуля/функции и возвращать исключительно значения 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]]) – Словарь типов аргументов для kwargs.
  • normalize_to_only_use_kwargs (bool) – Нужно ли нормализовать для использования только kwargs.
Возвращаемое значение:

Возвращает кортеж 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>>) [source]

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

Параметры:
  • replace_with (Узел) – Узел, на который нужно заменить все использования self.
  • delete_user_cb (Callable) – Обратный вызов, который вызывается для определения, нужно ли удалять данный пользователь узла self.
Возвращаемое значение:

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

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

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

Примечание

Обратная совместимость для этого 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, который был записан во время трассировки, если таковой имеется. Это свойство обычно заполняется Tracer.create_proxy. Чтобы записывать стеки вызовов во время трассировки для отладки, установите record_stack_traces = True на экземпляре Tracer.

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 (Вызываемый объект) – Метод forward() Module для вызова
  • args (Кортеж) – args вызова модуля
  • kwargs (Словарь) – kwargs вызова модуля
Возвращаемое значение:

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

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

Any

Примечание

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

create_arg(a) [source]

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

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

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

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

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

Параметры:

a (Любой тип) – Значение, которое должно быть сгенерировано как Argument в Graph.

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

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

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

Optional[Union[Tuple[Любой тип, …], Список[Любой тип], Словарь[str, Любой тип], slice, Узел, str, int, float, bool, complex, Тип данных, Тензор, Устройство, Формат памяти, Макет]]

Примечание

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

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

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

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

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

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

Вставляет узел графа, заданный целевым значением, аргументами, ключевыми аргументами и именем.

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

Примечание

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

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

Узел

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

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

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

Примечание

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

getattr(attr, attr_val, parameter_proxy_cache) [source]

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

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

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

Параметры:
  • attr (str) – Имя запрашиваемого атрибута
  • attr_val (Любой тип) – Значение атрибута
  • parametr_proxy_cache (Словарь[str, Любой тип]) – Кэширование имён атрибутов для прокси
Возвращаемое значение:

Возвращаемое значение вызова 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 (строка) – Module для получения квалифицированного имени.

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

строка

Примечание

Обратная совместимость для этого 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[Модуль, Вызов]) – Либо Module , либо функция, которую нужно отследить. Обратная совместимость для этого параметра гарантирована.
  • concrete_args (Optional[Dict[строка, любой]]) – Конкретные аргументы, которые не должны обрабатываться как прокси. Этот параметр экспериментальный, и его обратная совместимость НЕ гарантирована.
Возвращает:

Graph , представляющее семантику переданного 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_allclose(result, torch.neg(input).sigmoid())
Параметры:
  • module (GraphModule) – Модуль, который нужно выполнить
  • garbage_collect_values (bool) – Удалять ли значения после их последнего использования в рамках выполнения Модуля. Это гарантирует оптимальное использование памяти во время выполнения. Это можно отключить, например, для проверки всех промежуточных значений в выполнении, посмотрев на атрибут Interpreter.env.

Примечание

Обратная совместимость для этого 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[Кортеж[Любой, …], Список[Любой], Словарь[строка, Любой], срезы, Узел, строка, целое, вещественное, булево, комплексное, тип_данных, Тензор, устройство, форма_памяти, макет]]

Примечание

Обратная совместимость для этого 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 (булево) – Если 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_allclose(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 (Callable) — Подграф для сопоставления в gm для замены
  • replacement (Callable) — Подграф для замены 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, в то время как последним узлом является то, что возвращает функция.)

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

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

на

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

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

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

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

Примечание

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

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

Spec-Zone.ru

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