torch.func.jacrev
-
torch.func.jacrev(func, argnums=0, *, has_aux=False, chunk_size=None, _preallocate_and_copy=False)[source] -
Вычисляет якобиан
funcпо аргументу(ам) с индексомargnum, используя автоматическое дифференцирование в обратном режимеПримечание
Использование
chunk_size=1эквивалентно вычислению якобиана построчно с помощью цикла for, то есть ограниченияvmap()не применяются.- Параметры:
-
- func (function) – Функция Python, принимающая один или несколько аргументов, один из которых должен быть Tensor, и возвращающая один или несколько Tensor
- argnums (int or tuple[int, ...]) – Необязательный параметр: целое число или кортеж целых чисел, указывающий, по каким аргументам вычислять якобиан. Значение по умолчанию: 0.
-
has_aux (bool) – Флаг, указывающий, что
funcвозвращает кортеж(output, aux), в котором первый элемент — результат функции, для которой выполняется дифференцирование, а второй — вспомогательные объекты, для которых дифференцирование выполняться не будет. Значение по умолчанию: False. -
chunk_size (None or int) – Если значение равно None (по умолчанию), используется максимальный размер блока (эквивалентно одному вызову vmap над vjp для вычисления якобиана). Если значение равно 1, якобиан вычисляется построчно с помощью цикла for. Если значение не равно None, якобиан вычисляется по
chunk_sizeстрок за раз (эквивалентно нескольким вызовам vmap над vjp). Если при вычислении якобиана возникают проблемы с памятью, попробуйте указать ненулевой chunk_size.
- Возвращает:
-
Возвращает функцию, принимающую те же входные данные, что и
func, и возвращающую якобианfuncпо аргументу(ам) с индексомargnums. Еслиhas_aux is True, возвращённая функция вместо этого возвращает кортеж(jacobian, aux), гдеjacobian— якобиан, аaux— вспомогательные объекты, возвращённыеfunc. - Тип возвращаемого значения:
-
Callable[…, Any]
При поэлементной унарной операции якобиан будет диагональной матрицей
>>> 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.
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.func.jacrev.html