Spec-Zone.ru › PyTorch 2

torch.func.jacfwd

torch.func.jacfwd(func, argnums=0, has_aux=False, *, randomness='error')

Вычисляет якобиан func относительно аргумента(ов) с индексом argnum с использованием дифференцирования вперёд.

Параметры
  • func (функция) – Функция Python, принимающая один или несколько аргументов, один из которых должен быть тензором, и возвращающая один или несколько тензоров
  • argnums (int или Кортеж[int]) – Необязательный целочисленный или кортеж целых чисел, указывающий, по отношению к каким аргументам вычислять якобиан. По умолчанию: 0.
  • has_aux (bool) – Флаг, указывающий, что func возвращает кортеж (output, aux), где первый элемент — результат функции, подлежащей дифференцированию, а второй элемент — вспомогательные объекты, которые не будут дифференцироваться. По умолчанию: False.
  • randomness (str) – Флаг, указывающий, какой тип случайности использовать. Подробнее см. в vmap(). Допустимые значения: “different”, “same”, “error”. По умолчанию: “error”
Возвращает

Возвращает функцию, которая принимает те же входные данные, что и func, и возвращает якобиан func относительно аргумента(ов) с индексом argnums. Если has_aux is True, то возвращаемая функция вместо этого возвращает кортеж (jacobian, aux), где jacobian — якобиан, а aux — вспомогательные объекты, возвращаемые func.

Примечание

Возможно, эта API выдаст ошибку «дифференцирование вперёд не реализовано для оператора X». В таком случае, пожалуйста, отправьте отчёт об ошибке, и мы его приоритезируем. Альтернативный вариант — использовать jacrev(), у которого более полная поддержка операторов.

Базовое использование с точечной, унарной операцией даст диагональную матрицу в качестве якобиана

>>> from torch.func import jacfwd
>>> x = torch.randn(5)
>>> jacobian = jacfwd(torch.sin)(x)
>>> expected = torch.diag(torch.cos(x))
>>> assert torch.allclose(jacobian, expected)

jacfwd() можно комбинировать с vmap для получения пакетных якобианов:

>>> from torch.func import jacfwd, vmap
>>> x = torch.randn(64, 5)
>>> jacobian = vmap(jacfwd(torch.sin))(x)
>>> assert jacobian.shape == (64, 5, 5)

Если вы хотите вычислить и значение функции, и якобиан функции, используйте флаг has_aux для возвращения значения функции как вспомогательного объекта:

>>> from torch.func import jacfwd
>>> x = torch.randn(5)
>>>
>>> def f(x):
>>>   return x.sin()
>>>
>>> def g(x):
>>>   result = f(x)
>>>   return result, result
>>>
>>> jacobian_f, f_x = jacfwd(g, has_aux=True)(x)
>>> assert torch.allclose(f_x, f(x))

Кроме того, jacrev() можно комбинировать с собой или с jacrev() для получения матриц Гессе

>>> from torch.func import jacfwd, jacrev
>>> def f(x):
>>>   return x.sin().sum()
>>>
>>> x = torch.randn(5)
>>> hessian = jacfwd(jacrev(f))(x)
>>> assert torch.allclose(hessian, torch.diag(-x.sin()))

По умолчанию jacfwd() вычисляет якобиан относительно первого входного значения. Однако можно вычислить якобиан относительно другого аргумента, используя argnums:

>>> from torch.func import jacfwd
>>> def f(x, y):
>>>   return x + y ** 2
>>>
>>> x, y = torch.randn(5), torch.randn(5)
>>> jacobian = jacfwd(f, argnums=1)(x, y)
>>> expected = torch.diag(2 * y)
>>> assert torch.allclose(jacobian, expected)

Кроме того, передача кортежа в argnums позволит вычислить якобиан по отношению к нескольким аргументам

>>> from torch.func import jacfwd
>>> def f(x, y):
>>>   return x + y ** 2
>>>
>>> x, y = torch.randn(5), torch.randn(5)
>>> jacobian = jacfwd(f, argnums=(0, 1))(x, y)
>>> expectedX = torch.diag(torch.ones_like(x))
>>> expectedY = torch.diag(2 * y)
>>> assert torch.allclose(jacobian[0], expectedX)
>>> assert torch.allclose(jacobian[1], expectedY)

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

Spec-Zone.ru

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