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