Spec-Zone.ru › PyTorch 2

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

Spec-Zone.ru

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