Краткий обзор torch.func
Создано: 12 июня 2025 г. | Последнее обновление: 12 июня 2025 г.
Что такое torch.func?
torch.func, ранее известная как functorch, — это библиотека для компонуемых преобразований функций в PyTorch, подобных JAX.
- «Преобразование функции» — это функция высшего порядка, которая принимает числовую функцию и возвращает новую функцию, вычисляющую другую величину.
- 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() накладывает ограничения на код, к которому его можно применять. Подробнее см. в разделе Ограничения UX.
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)
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/func.whirlwind_tour.html