Spec-Zone.ru › PyTorch 2.14

torch.func.vjp

torch.func.vjp(func, *primals, has_aux=False) [source]

Сокращение от vector-Jacobian product (произведение вектора на матрицу Якоби): возвращает кортеж, содержащий результаты применения func к primals, и функцию, которая при получении cotangents вычисляет якобиан func в режиме обратного распространения по отношению к primals, умноженный на cotangents.

Параметры:
  • func (Callable[..., Any]) – Функция Python, принимающая один или несколько аргументов. Должна возвращать один или несколько тензоров.
  • primals (Tensors) – Позиционные аргументы для func, каждый из которых должен быть тензором. Возвращаемая функция также вычисляет производную по этим аргументам
  • has_aux (bool) – Флаг, указывающий, что func возвращает кортеж (output, aux), в котором первый элемент — результат функции, для которой вычисляется производная, а второй — другие вспомогательные объекты, для которых производная вычисляться не будет. Значение по умолчанию: False.
Возвращает:

Возвращает кортеж (output, vjp_fn), содержащий результат применения func к primals, и функцию, которая вычисляет vjp для func по всем primals, используя котангенты, переданные возвращаемой функции. Если has_aux is True, возвращается кортеж (output, vjp_fn, aux). Возвращаемая функция vjp_fn возвращает кортеж, содержащий каждый VJP.

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

tuple[Any, Callable[…, Any]] | tuple[Any, Callable[…, Any], Any]

В простых случаях vjp() работает так же, как grad()

>>> x = torch.randn([5])
>>> f = lambda x: x.sin().sum()
>>> (_, vjpfunc) = torch.func.vjp(f, x)
>>> grad = vjpfunc(torch.tensor(1.0))[0]
>>> assert torch.allclose(grad, torch.func.grad(f)(x))

Однако vjp() поддерживает функции с несколькими выходами, если передать котангенты для каждого из выходов

>>> x = torch.randn([5])
>>> f = lambda x: (x.sin(), x.cos())
>>> (_, vjpfunc) = torch.func.vjp(f, x)
>>> vjps = vjpfunc((torch.ones([5]), torch.ones([5])))
>>> assert torch.allclose(vjps[0], x.cos() + -x.sin())

vjp() также поддерживает выходные данные в виде структур Python

>>> x = torch.randn([5])
>>> f = lambda x: {"first": x.sin(), "second": x.cos()}
>>> (_, vjpfunc) = torch.func.vjp(f, x)
>>> cotangents = {"first": torch.ones([5]), "second": torch.ones([5])}
>>> vjps = vjpfunc(cotangents)
>>> assert torch.allclose(vjps[0], x.cos() + -x.sin())

Функция, возвращаемая vjp(), вычисляет частные производные по каждому из primals

>>> x, y = torch.randn([5, 4]), torch.randn([4, 5])
>>> (_, vjpfunc) = torch.func.vjp(torch.matmul, x, y)
>>> cotangents = torch.randn([5, 5])
>>> vjps = vjpfunc(cotangents)
>>> assert len(vjps) == 2
>>> assert torch.allclose(vjps[0], torch.matmul(cotangents, y.transpose(0, 1)))
>>> assert torch.allclose(vjps[1], torch.matmul(x.transpose(0, 1), cotangents))

primals — это позиционные аргументы для f. Для всех kwargs используются значения по умолчанию

>>> x = torch.randn([5])
>>> def f(x, scale=4.):
>>>   return x * scale
>>>
>>> (_, vjpfunc) = torch.func.vjp(f, x)
>>> vjps = vjpfunc(torch.ones_like(x))
>>> assert torch.allclose(vjps[0], torch.full(x.shape, 4.0))

Примечание

Использование PyTorch torch.no_grad вместе с vjp. Случай 1: использование torch.no_grad внутри функции:

>>> def f(x):
>>>     with torch.no_grad():
>>>         c = x ** 2
>>>     return x - c

В этом случае vjp(f)(x) будет учитывать внутренний torch.no_grad.

Случай 2: использование vjp внутри менеджера контекста torch.no_grad:

>>> with torch.no_grad():
>>>     vjp(f)(x)

В этом случае vjp будет учитывать внутренний torch.no_grad, но не внешний. Это связано с тем, что vjp — это «преобразование функции»: результат не должен зависеть от результата менеджера контекста, находящегося за пределами f.

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

Spec-Zone.ru

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