torch.func.hessian
-
torch.func.hessian(func, argnums=0)[исходный код] -
Вычисляет матрицу Гессе для
funcпо аргументам с индексомargnumс помощью стратегии «прямой режим поверх обратного».Стратегия «прямой режим поверх обратного» (композиция
jacfwd(jacrev(func))) — хороший вариант по умолчанию для высокой производительности. Матрицы Гессе можно вычислять и с помощью других комбинацийjacfwd()иjacrev(), напримерjacfwd(jacfwd(func))илиjacrev(jacrev(func)).- Параметры:
-
- func (function) – Функция Python, принимающая один или несколько аргументов, по крайней мере один из которых должен быть тензором, и возвращающая один или несколько тензоров
- argnums (int or tuple[int, ...]) – Необязательный параметр: целое число или кортеж целых чисел, указывающий, для каких аргументов нужно получить матрицу Гессе. Значение по умолчанию: 0.
- Возвращает:
-
Возвращает функцию, принимающую те же входные данные, что и
func, и возвращающую матрицу Гессе дляfuncпо аргументам с индексомargnums. - Тип возвращаемого значения:
-
Callable[…, Any]
Примечание
При использовании этого API может возникнуть ошибка «forward-mode AD not implemented for operator X» («прямой режим автоматического дифференцирования не реализован для оператора X»). В этом случае сообщите об ошибке — мы уделим ей приоритетное внимание. В качестве альтернативы можно использовать
jacrev(jacrev(func)), который поддерживает больше операторов.Для функции R^N -> R^1 простое применение возвращает матрицу Гессе размера N x N:
>>> from torch.func import hessian >>> def f(x): >>> return x.sin().sum() >>> >>> x = torch.randn(5) >>> hess = hessian(f)(x) # equivalent to jacfwd(jacrev(f))(x) >>> assert torch.allclose(hess, torch.diag(-x.sin()))
© 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.hessian.html