Spec-Zone.ru › PyTorch 2.14

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

Spec-Zone.ru

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