torch.func.jvp
-
torch.func.jvp(func, primals, tangents, *, strict=False, has_aux=False) -
Обозначая произведение Якобиана и вектора, возвращает кортеж, содержащий результат
func(*primals)и «Якобианfunc, вычисленный вprimals» умноженный наtangents. Это также известно как автодифференцирование по направлению вперёд.- Параметры
-
- func (функция) – Функция Python, принимающая один или несколько аргументов, один из которых должен быть тензором, и возвращающая один или несколько тензоров
-
primals (тензоры) – Позиционные аргументы для
func, все из которых должны быть тензорами. Возвращаемая функция также будет вычислять производную по этим аргументам -
tangents (тензоры) – «Вектор», для которого вычисляется произведение Якобиана и вектора. Должен иметь ту же структуру и размеры, что и входные данные для
func. -
has_aux (bool) – Флаг, указывающий, что
funcвозвращает кортеж(output, aux), где первый элемент — результат функции, подлежащей дифференцированию, а второй элемент — другие вспомогательные объекты, которые не будут дифференцироваться. По умолчанию: False.
- Возвращает
-
Возвращает кортеж
(output, jvp_out)содержащий результатfunc, вычисленный вprimalsи произведение Якобиана и вектора. Еслиhas_aux is True, то вместо этого возвращается кортеж(output, jvp_out, aux).
Примечание
Вы можете увидеть ошибку API с сообщением «автодифференцирование по направлению вперёд не реализовано для оператора X». В таком случае, пожалуйста, подайте отчёт об ошибке, и мы его рассмотрим в приоритетном порядке.
jvp полезно, когда вы хотите вычислить градиенты функции R^1 -> R^N
>>> from torch.func import jvp >>> x = torch.randn([]) >>> f = lambda x: x * torch.tensor([1., 2., 3]) >>> value, grad = jvp(f, (x,), (torch.tensor(1.),)) >>> assert torch.allclose(value, f(x)) >>> assert torch.allclose(grad, torch.tensor([1., 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)
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.func.jvp.html