Пользовательские бэкэнды
Обзор
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