Spec-Zone.ru › PyTorch 2

torch.func.jacrev

torch.func.jacrev(func, argnums=0, *, has_aux=False, chunk_size=None, _preallocate_and_copy=False)

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

Примечание

Использование chunk_size=1 эквивалентно вычислению якобиана построчно с помощью цикла for, т.е. ограничения vmap() не применяются.

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

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

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

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

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

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

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

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

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

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

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

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

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

>>> from torch.func import jacrev
>>> def f(x, y):
>>>   return x + y ** 2
>>>
>>> x, y = torch.randn(5), torch.randn(5)
>>> jacobian = jacrev(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)

Примечание

Использование PyTorch torch.no_grad вместе с jacrev. Случай 1: Использование torch.no_grad внутри функции:

>>> def f(x):
>>>     with torch.no_grad():
>>>         c = x ** 2
>>>     return x - c

В этом случае jacrev(f)(x) будет учитывать внутренний torch.no_grad.

Случай 2: Использование jacrev внутри контекстного менеджера torch.no_grad:

>>> with torch.no_grad():
>>>     jacrev(f)(x)

В этом случае jacrev будет учитывать внутренний torch.no_grad, но не внешний. Это потому, что jacrev — это «трансформация функций»: её результат не должен зависеть от результата контекстного менеджера вне f.

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

Spec-Zone.ru

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