Spec-Zone.ru › PyTorch 2

Пользовательские бэкэнды

Обзор

torch.compile предоставляет простой способ для пользователей определять пользовательские бэкэнды.

Функция бэкэнда имеет контракт (gm: torch.fx.GraphModule, example_inputs: List[torch.Tensor]) -> Callable.

Функции бэкэндов могут вызываться TorchDynamo, компонентом трассировки графа в torch.compile, после трассировки графа FX и ожидается, что они вернут скомпилированную функцию, эквивалентную прослеженному графу FX. Возвращаемая вызываемая функция должна иметь тот же контракт, что и функция forward исходного torch.fx.GraphModule , переданного в бэкенд: (*args: torch.Tensor) -> List[torch.Tensor].

Для того, чтобы TorchDynamo мог вызвать ваш бэкенд, передайте вашу функцию бэкэнда в качестве аргумента backend в torch.compile. Например,

import torch

def my_custom_backend(gm, example_inputs):
    return gm.forward

def f(...):
    ...

f_opt = torch.compile(f, backend=my_custom_backend)

@torch.compile(backend=my_custom_backend)
def g(...):
    ...

Ниже приведены дополнительные примеры.

Регистрация пользовательских бэкэндов

Вы можете зарегистрировать свой бэкенд, используя декоратор register_backend , например,

from torch._dynamo.optimizations import register_backend

@register_backend
def my_compiler(gm, example_inputs):
    ...

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

Подсказка

Дополнительную информацию о entry_points можно найти в документации по упаковке пакетов Python.

Для регистрации вашего бэкэнда через entry_points, вы можете добавить свою функцию бэкэнда в группу точек входа torch_dynamo_backends в файле setup.py вашего пакета, как показано ниже:

...
setup(
    ...
    'torch_dynamo_backends': [
        'my_compiler = your_module.submodule:my_compiler',
    ]
    ...
)

Пожалуйста, замените my_compiler перед = на имя вашего бэкэнда и замените часть после = на имя модуля и функции вашей функции бэкэнда. Точка входа будет добавлена в вашу среду Python после установки пакета. Когда вы вызываете torch.compile(model, backend="my_compiler"), PyTorch сначала ищет бэкенд под именем my_compiler, зарегистрированный с помощью register_backend. Если он не найден, поиск продолжается во всех бэкендах, зарегистрированных с помощью entry_points.

Регистрация выполняет две задачи:

  • Вы можете передать строку с именем функции вашего бэкэнда в torch.compile вместо самой функции, например, torch.compile(model, backend="my_compiler").
  • Это требуется для использования с минификатором. Любой сгенерированный код из минификатора должен вызывать ваш код, регистрирующий вашу функцию бэкэнда, обычно с помощью оператора import.

Пользовательские бэкэнды после AOTAutograd

Возможна разработка пользовательских бэкэндов, вызываемых AOTAutograd, а не TorchDynamo. Это полезно по двум основным причинам:

  • Пользователи могут определить бэкэнды, поддерживающие обучение моделей, поскольку AOTAutograd может сгенерировать обратный граф для компиляции.
  • AOTAutograd генерирует графы FX, состоящие из канонических операций Aten. В результате пользовательские бэкэнды должны поддерживать только канонический набор операций Aten, что значительно меньше, чем весь набор операций torch/Aten.

Оборачивайте свой бэкенд с помощью torch._dynamo.optimizations.training.aot_autograd и используйте torch.compile с аргументом backend как и раньше. Функции бэкэнда, обернутые aot_autograd, должны иметь тот же контракт, что и раньше.

Функции бэкэндов передаются в aot_autograd через аргументы fw_compiler (компилятор прямого прохода) или bw_compiler (компилятор обратного прохода). Если bw_compiler не указано, функция обратной компиляции по умолчанию совпадает с функцией прямой компиляции.

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

Например,

from torch._dynamo.optimizations.training import aot_autograd
from functorch.compile import make_boxed_func

def my_compiler(gm, example_inputs):
    return make_boxed_func(gm.forward)

my_backend = aot_autograd(fw_compiler=my_compiler)  # bw_compiler=my_compiler

model_opt = torch.compile(model, backend=my_backend)

Примеры

Отладка бэкэнда

Если вы хотите лучше понять, что происходит во время компиляции, вы можете создать пользовательский компилятор, который в этом разделе называется бэкэндом, который будет выводить красиво отформатированный граф fx GraphModule, извлечённый из анализа байткода Dynamo, и возвращать вызываемую функцию forward().

Например:

from typing import List
import torch
def my_compiler(gm: torch.fx.GraphModule, example_inputs: List[torch.Tensor]):
    print("my_compiler() called with FX graph:")
    gm.graph.print_tabular()
    return gm.forward  # return a python callable
@torch.compile(backend=my_compiler)
def fn(x, y):
    a = torch.cos(x)
    b = torch.sin(y)
    return a + b
fn(torch.randn(10), torch.randn(10))

Выполнение приведенного выше примера приводит к следующему выводу:

my_compiler() called with FX graph:
opcode         name    target                                                  args        kwargs
-------------  ------  ------------------------------------------------------  ----------  --------
placeholder    x       x                                                       ()          {}
placeholder    y       y                                                       ()          {}
call_function  cos     <built-in method cos of type object at 0x7f1a894649a8>  (x,)        {}
call_function  sin     <built-in method sin of type object at 0x7f1a894649a8>  (y,)        {}
call_function  add     <built-in function add>                                 (cos, sin)  {}
output         output  output                                                  ((add,),)   {}

Это работает и для torch.nn.Module, как показано ниже:

from typing import List
import torch
def my_compiler(gm: torch.fx.GraphModule, example_inputs: List[torch.Tensor]):
    print("my_compiler() called with FX graph:")
    gm.graph.print_tabular()
    return gm.forward  # return a python callable
class MockModule(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.relu = torch.nn.ReLU()
    def forward(self, x):
        return self.relu(torch.cos(x))
mod = MockModule()
optimized_mod = torch.compile(mod, backend=my_compiler)
optimized_mod(torch.randn(10))

Давайте рассмотрим ещё один пример с потоком управления:

from typing import List
import torch
def my_compiler(gm: torch.fx.GraphModule, example_inputs: List[torch.Tensor]):
    print("my_compiler() called with FX graph:")
    gm.graph.print_tabular()
    return gm.forward  # return a python callable
@torch.compile(backend=my_compiler)
def toy_example(a, b):
    x = a / (torch.abs(a) + 1)
    if b.sum() < 0:
        b = b * -1
    return x * b
for _ in range(100):
    toy_example(torch.randn(10), torch.randn(10))

Выполнение этого примера приводит к следующему выводу:

my_compiler() called with FX graph:
opcode         name     target                                                  args              kwargs
-------------  -------  ------------------------------------------------------  ----------------  --------
placeholder    a        a                                                       ()                {}
placeholder    b        b                                                       ()                {}
call_function  abs_1    <built-in method abs of type object at 0x7f8d259298a0>  (a,)              {}
call_function  add      <built-in function add>                                 (abs_1, 1)        {}
call_function  truediv  <built-in function truediv>                             (a, add)          {}
call_method    sum_1    sum                                                     (b,)              {}
call_function  lt       <built-in function lt>                                  (sum_1, 0)        {}
output         output   output                                                  ((truediv, lt),)  {}

my_compiler() called with FX graph:
opcode         name    target                   args         kwargs
-------------  ------  -----------------------  -----------  --------
placeholder    b       b                        ()           {}
placeholder    x       x                        ()           {}
call_function  mul     <built-in function mul>  (b, -1)      {}
call_function  mul_1   <built-in function mul>  (x, mul)     {}
output         output  output                   ((mul_1,),)  {}

my_compiler() called with FX graph:
opcode         name    target                   args       kwargs
-------------  ------  -----------------------  ---------  --------
placeholder    b       b                        ()         {}
placeholder    x       x                        ()         {}
call_function  mul     <built-in function mul>  (x, b)     {}
output         output  output                   ((mul,),)  {}

Порядок последних двух графов является не детерминированным, в зависимости от того, какой из них встречается компилятором just-in-time первым.

Быстрый бэкенд

Интеграция пользовательского бэкэнда, обеспечивающего превосходную производительность, также проста, и мы интегрируем реальный бэкенд с optimize_for_inference:

def optimize_for_inference_compiler(gm: torch.fx.GraphModule, example_inputs: List[torch.Tensor]):
    scripted = torch.jit.script(gm)
    return torch.jit.optimize_for_inference(scripted)

И тогда вы должны иметь возможность оптимизировать любой существующий код с помощью:

@torch.compile(backend=optimize_for_inference_compiler)
def code_to_accelerate():
    ...

Комбинируемые бэкэнды

TorchDynamo включает в себя множество бэкэндов, которые можно найти в backends.py или torch._dynamo.list_backends(). Вы можете объединить эти бэкэнды вместе с помощью следующего кода:

from torch._dynamo.optimizations import BACKENDS
 def my_compiler(gm: torch.fx.GraphModule, example_inputs: List[torch.Tensor]):
     try:
         trt_compiled = BACKENDS["tensorrt"](gm, example_inputs)
         if trt_compiled is not None:
             return trt_compiled
     except Exception:
         pass
     # first backend failed, try something else...
     try:
         inductor_compiled = BACKENDS["inductor"](gm, example_inputs)
         if inductor_compiled is not None:
             return inductor_compiled
     except Exception:
         pass
     return gm.forward

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

Spec-Zone.ru

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