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