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). - Тип возвращаемого значения:
Примечание
При использовании этого 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