Spec-Zone.ru › PyTorch 2.14

torch.func.jvp

torch.func.jvp(func, primals, tangents, *, strict=False, has_aux=False) [исходный код]

Обозначая произведение Якобиана на вектор, возвращает кортеж, содержащий результат func(*primals) и «Якобиан func, вычисленный в точке primals», умноженный на tangents. Это также известно как автоматическое дифференцирование в прямом режиме.

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

Возвращает кортеж (output, jvp_out), содержащий результат func, вычисленный в точке primals, и произведение Якобиана на вектор. Если has_aux is True, вместо этого возвращает кортеж (output, jvp_out, aux).

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

tuple[Any, Any] | tuple[Any, Any, Any]

Примечание

При использовании этого API может возникнуть ошибка «forward-mode AD not implemented for operator X». В этом случае отправьте сообщение об ошибке, и мы займёмся её устранением в приоритетном порядке.

jvp полезен, если требуется вычислить градиенты функции R^1 -> R^N

>>> from torch.func import jvp
>>> x = torch.randn([])
>>> f = lambda x: x * torch.tensor([1.0, 2.0, 3])
>>> warnings.filterwarnings(
...     "ignore", message=".*torch.jit.script"
... )  # docs: hide
>>> value, grad = jvp(f, (x,), (torch.tensor(1.0),))
>>> assert torch.allclose(value, f(x))
>>> assert torch.allclose(grad, torch.tensor([1.0, 2, 3]))

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

>>> from torch.func import jvp
>>> x = torch.randn(5)
>>> y = torch.randn(5)
>>> f = lambda x, y: (x * y)
>>> _, output = jvp(f, (x, y), (torch.ones(5), torch.ones(5)))
>>> assert torch.allclose(output, x + y)

© 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.jvp.html

Spec-Zone.ru

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