Обзор torch.func
Что такое torch.func?
torch.func, ранее известный как functorch, — это библиотека для создания композируемых преобразований функций в стиле JAX для PyTorch.
- «Преобразование функции» — это функция высшего порядка, которая принимает численную функцию и возвращает новую функцию, вычисляющую другое значение.
- torch.func имеет преобразования для автоматического дифференцирования (
grad(f)возвращает функцию, вычисляющую градиентf), преобразования для векторизации/пакетирования (vmap(f)возвращает функцию, вычисляющуюfпо пакетам входных данных) и другие. - Эти преобразования функций можно произвольно комбинировать. Например, комбинируя
vmap(grad(f))вычисляется значение, называемое градиентами по образцу, которое PyTorch пока не может эффективно вычислить.
Зачем нужны композируемые преобразования функций?
Существует ряд задач, которые трудно решить в PyTorch: - вычисление градиентов по образцу (или других значений по образцу)
- запуск ансамблей моделей на одном компьютере
- эффективное объединение задач во внутреннем цикле MAML
- эффективное вычисление Якобианов и Гессианов
- эффективное вычисление объединенных Якобианов и Гессианов
Комбинирование преобразований vmap(), grad(), vjp() и jvp() позволяет нам выразить вышеперечисленное, не разрабатывая отдельную подсистему для каждого случая.
Какие преобразования доступны?
grad() (вычисление градиента)
grad(func) — это наше преобразование для вычисления градиента. Оно возвращает новую функцию, которая вычисляет градиенты func. Предполагается, что func возвращает тензор с одним элементом, и по умолчанию оно вычисляет градиенты выходного значения func по отношению к первому входу.
import torch from torch.func import grad x = torch.randn([]) cos_x = grad(lambda x: torch.sin(x))(x) assert torch.allclose(cos_x, x.cos()) # Second-order gradients neg_sin_x = grad(grad(lambda x: torch.sin(x)))(x) assert torch.allclose(neg_sin_x, -x.sin())
vmap() (автоматическая векторизация)
Примечание: vmap() накладывает ограничения на код, к которому можно применить это преобразование. Более подробная информация доступна в разделе Ограничения.
vmap(func)(*inputs) — это преобразование, которое добавляет размерность ко всем операциям с тензорами в func. vmap(func) возвращает новую функцию, которая применяет func к некоторой размерности (по умолчанию 0) каждого тензора во входных данных.
vmap полезен для скрытия размерности пакетов: можно написать функцию func, которая работает с примерами, а затем поднять ее до функции, которая может принимать пакеты примеров с помощью vmap(func), что упрощает работу с моделями:
import torch
from torch.func import vmap
batch_size, feature_size = 3, 5
weights = torch.randn(feature_size, requires_grad=True)
def model(feature_vec):
# Very simple linear model with activation
assert feature_vec.dim() == 1
return feature_vec.dot(weights).relu()
examples = torch.randn(batch_size, feature_size)
result = vmap(model)(examples)
При композиции с grad(), vmap() можно использовать для вычисления градиентов по образцу:
from torch.func import vmap
batch_size, feature_size = 3, 5
def model(weights,feature_vec):
# Very simple linear model with activation
assert feature_vec.dim() == 1
return feature_vec.dot(weights).relu()
def compute_loss(weights, example, target):
y = model(weights, example)
return ((y - target) ** 2).mean() # MSELoss
weights = torch.randn(feature_size, requires_grad=True)
examples = torch.randn(batch_size, feature_size)
targets = torch.randn(batch_size)
inputs = (weights,examples, targets)
grad_weight_per_example = vmap(grad(compute_loss), in_dims=(None, 0, 0))(*inputs)
vjp() (векторно-Якобианное произведение)
Преобразование vjp() применяет func к inputs и возвращает новую функцию, вычисляющую векторно-Якобианное произведение (vjp) для заданных cotangents тензоров.
from torch.func import vjp inputs = torch.randn(3) func = torch.sin cotangents = (torch.randn(3),) outputs, vjp_fn = vjp(func, inputs); vjps = vjp_fn(*cotangents)
jvp() (Якобианно-векторное произведение)
Преобразование jvp() вычисляет Якобианно-векторные произведения и также известно как «дифференцирование в прямом режиме». В отличие от большинства других преобразований, это не функция высшего порядка, но оно возвращает выходные значения func(inputs) и jvp.
from torch.func import jvp x = torch.randn(5) y = torch.randn(5) f = lambda x, y: (x * y) _, out_tangent = jvp(f, (x, y), (torch.ones(5), torch.ones(5))) assert torch.allclose(out_tangent, x + y)
jacrev(), jacfwd() и hessian()
Преобразование jacrev() возвращает новую функцию, которая принимает x и возвращает Якобиан функции относительно x с использованием дифференцирования в обратном режиме.
from torch.func import jacrev x = torch.randn(5) jacobian = jacrev(torch.sin)(x) expected = torch.diag(torch.cos(x)) assert torch.allclose(jacobian, expected)
jacrev() можно объединить с vmap() для получения объединенных Якобианов:
x = torch.randn(64, 5) jacobian = vmap(jacrev(torch.sin))(x) assert jacobian.shape == (64, 5, 5)
jacfwd() — это прямая замена для jacrev, которая вычисляет Якобианы с использованием дифференцирования в прямом режиме:
from torch.func import jacfwd x = torch.randn(5) jacobian = jacfwd(torch.sin)(x) expected = torch.diag(torch.cos(x)) assert torch.allclose(jacobian, expected)
Объединение jacrev() с самим собой или jacfwd() может привести к вычислению Гессиана:
def f(x):
return x.sin().sum()
x = torch.randn(5)
hessian0 = jacrev(jacrev(f))(x)
hessian1 = jacfwd(jacrev(f))(x)
hessian() — это вспомогательная функция, которая объединяет jacfwd и jacrev:
from torch.func import hessian
def f(x):
return x.sin().sum()
x = torch.randn(5)
hess = hessian(f)(x)
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/func.whirlwind_tour.html