Spec-Zone.ru › PyTorch 2.14

Краткий обзор 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

Spec-Zone.ru

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