Spec-Zone.ru › PyTorch 2

torch.func.functionalize

torch.func.functionalize(func, *, remove='mutations')

functionalize — это преобразование, которое можно использовать для удаления (промежуточных) мутаций и алиасов из функции, сохраняя при этом семантику функции.

functionalize(func) возвращает новую функцию с той же семантикой, что и func, но со всеми промежуточными мутациями, удалёнными. Каждая операция in-place, выполняемая над промежуточным тензором: intermediate.foo_() заменяется её эквивалентом out-of-place: intermediate_updated = intermediate.foo().

functionalize полезно для отправки программы PyTorch на бэкэнды или компиляторы, которые не могут легко представить мутации или операторы алиасирования.

Параметры
  • func (Callable) – Функция Python, принимающая один или несколько аргументов.
  • remove (str) – Необязательный строковый аргумент, принимающий значения ‘mutations’ или ‘mutations_and_views’. Если передано ‘mutations’, все мутирующие операторы будут заменены на свои немутирующие эквиваленты. Если передано ‘mutations_and_views’, дополнительно все операторы алиасирования будут заменены на свои не-алиасные эквиваленты. По умолчанию: ‘mutations’.
Возвращаемое значение

Возвращает новую «функционализированную» функцию. Она принимает те же входные данные, что и func, и имеет такое же поведение, но любые мутации (и необязательно алиасирования), выполняемые над промежуточными тензорами в функции, будут удалены.

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

Callable

functionalize также удалит мутации (и представления), которые были выполнены над входными данными функции. Однако для сохранения семантики functionalize «исправит» мутации после завершения преобразования, обнаружив, должны ли были быть изменены какие-либо входные тензоры и скопировав новые данные обратно в входные данные, если это необходимо.

Пример:

>>> import torch
>>> from torch.fx.experimental.proxy_tensor import make_fx
>>> from torch.func import functionalize
>>>
>>> # A function that uses mutations and views, but only on intermediate tensors.
>>> def f(a):
...     b = a + 1
...     c = b.view(-1)
...     c.add_(1)
...     return b
...
>>> inpt = torch.randn(2)
>>>
>>> out1 = f(inpt)
>>> out2 = functionalize(f)(inpt)
>>>
>>> # semantics are the same (outputs are equivalent)
>>> print(torch.allclose(out1, out2))
True
>>>
>>> f_traced = make_fx(f)(inpt)
>>> f_no_mutations_traced = make_fx(functionalize(f))(inpt)
>>> f_no_mutations_and_views_traced = make_fx(functionalize(f, remove='mutations_and_views'))(inpt)
>>>
>>> print(f_traced.code)



def forward(self, a_1):
    add = torch.ops.aten.add(a_1, 1);  a_1 = None
    view = torch.ops.aten.view(add, [-1])
    add_ = torch.ops.aten.add_(view, 1);  view = None
    return add

>>> print(f_no_mutations_traced.code)



def forward(self, a_1):
    add = torch.ops.aten.add(a_1, 1);  a_1 = None
    view = torch.ops.aten.view(add, [-1]);  add = None
    add_1 = torch.ops.aten.add(view, 1);  view = None
    view_1 = torch.ops.aten.view(add_1, [2]);  add_1 = None
    return view_1

>>> print(f_no_mutations_and_views_traced.code)



def forward(self, a_1):
    add = torch.ops.aten.add(a_1, 1);  a_1 = None
    view_copy = torch.ops.aten.view_copy(add, [-1]);  add = None
    add_1 = torch.ops.aten.add(view_copy, 1);  view_copy = None
    view_copy_1 = torch.ops.aten.view_copy(add_1, [2]);  add_1 = None
    return view_copy_1


>>> # A function that mutates its input tensor
>>> def f(a):
...     b = a.view(-1)
...     b.add_(1)
...     return a
...
>>> f_no_mutations_and_views_traced = make_fx(functionalize(f, remove='mutations_and_views'))(inpt)
>>> #
>>> # All mutations and views have been removed,
>>> # but there is an extra copy_ in the graph to correctly apply the mutation to the input
>>> # after the function has completed.
>>> print(f_no_mutations_and_views_traced.code)



def forward(self, a_1):
    view_copy = torch.ops.aten.view_copy(a_1, [-1])
    add = torch.ops.aten.add(view_copy, 1);  view_copy = None
    view_copy_1 = torch.ops.aten.view_copy(add, [2]);  add = None
    copy_ = torch.ops.aten.copy_(a_1, view_copy_1);  a_1 = None
    return view_copy_1
Существует несколько «режимов сбоя» для functionalize, на которые стоит обратить внимание:
  1. Как и другие преобразования torch.func, functionalize() не работает с функциями, которые напрямую используют .backward(). То же самое верно для torch.autograd.grad. Если вы хотите использовать autograd, вы можете вычислять градиенты непосредственно с помощью functionalize(grad(f)).
  2. Как и другие преобразования torch.func, functionalize() не работает с глобальным состоянием. Если вы вызываете functionalize(f) для функции, которая использует представления / мутации нелокального состояния, функционализация просто ничего не сделает и передаст вызовы представления/мутации напрямую в бэкэнд. Один из способов обойти это — убедиться, что любое создание нелокального состояния обернуто в более крупную функцию, для которой вы затем вызываете functionalize.
  3. resize_() имеет некоторые ограничения: функционализация будет работать только с программами, которые используют `resize_()`, при условии, что тензор, который изменяется, не является представлением.
  4. as_strided() имеет некоторые ограничения: функционализация не будет работать с вызовами as_strided(), которые приводят к тензорам с перекрывающейся памятью.

Наконец, полезная модель для понимания функционализации заключается в том, что большинство пользовательских программ PyTorch написаны с использованием публичного API Torch. При выполнении операторы Torch обычно декомпозируются в наш внутренний C++ API «ATen». Логика функционализации полностью происходит на уровне ATen. Функционализация знает, как взять каждый оператор алиасирования в ATen и сопоставить его с его не-алиасинговым эквивалентом (например, tensor.view({-1}) -> at::view_copy(tensor, {-1}) ) и как взять каждый мутирующий оператор в ATen и сопоставить его с его не-мутирующим эквивалентом (например, tensor.add_(1) -> at::add(tensor, -1) ), одновременно отслеживая алиасы и мутации, чтобы знать, когда необходимо что-то исправить. Информация о том, какие операторы ATen являются алиасинговыми или мутирующими, берется из https://github.com/pytorch/pytorch/blob/master/aten/src/ATen/native/native_functions.yaml.

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

Spec-Zone.ru

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