Функция gradgradcheck
Функция gradgradcheck проверяет корректность вычисления вторых производных
пользовательских функций относительно входных данных. Она является расширением
функции gradcheck и предназначена для проверки функций, у которых
необходимо вычислить градиенты второго порядка. В параметрах функции передаются
проверяемая функция, входные тензоры, а также дополнительные аргументы,
управляющие процессом проверки.
Синтаксис
torch.autograd.gradgradcheck(
func,
inputs,
grad_outputs=None,
eps=1e-6,
atol=1e-5,
rtol=1e-3,
check_undefined_grad=True,
check_grad_grad_types=False
)
Основные параметры
func - проверяемая функция, принимающая входные тензоры
inputs - кортеж входных тензоров
grad_outputs - тензоры градиентов на выходе (по умолчанию None)
eps - шаг для численного дифференцирования (по умолчанию 1e-6)
atol - абсолютная допустимая погрешность (по умолчанию 1e-5)
rtol - относительная допустимая погрешность (по умолчанию 1e-3)
check_undefined_grad - проверять ли неопределённые градиенты
check_grad_grad_types - проверять ли типы градиентов градиентов
Возвращаемое значение: булево значение True, если проверка пройдена,
иначе выбрасывается исключение RuntimeError с описанием ошибки.
Пример проверки функции с градиентами второго порядка
Создадим простую функцию, которая вычисляет квадрат суммы входных тензоров,
и проверим вторые производные с помощью gradgradcheck:
import torch
from torch.autograd import gradgradcheck
def func(x, y):
return (x + y) ** 2
torch.manual_seed(0)
x = torch.randn(3, requires_grad=True)
y = torch.randn(3, requires_grad=True)
res = gradgradcheck(func, (x, y))
print(res)
Результат выполнения кода:
True
Пример с передачей grad_outputs
При необходимости можно передать конкретные градиенты на выходе
через параметр grad_outputs:
import torch
from torch.autograd import gradgradcheck
def func(x):
return x ** 3
torch.manual_seed(1)
x = torch.randn(2, requires_grad=True)
grad_out = torch.ones(2)
res = gradgradcheck(func, (x,), grad_outputs=(grad_out,))
print(res)
Результат выполнения кода:
True
Пример проверки функции с несколькими выходами
Для функций, возвращающих несколько значений, проверка вторых производных также работает корректно:
import torch
from torch.autograd import gradgradcheck
def func(x, y):
return x ** 2 + y ** 2, x * y
torch.manual_seed(2)
x = torch.randn(2, requires_grad=True)
y = torch.randn(2, requires_grad=True)
res = gradgradcheck(func, (x, y))
print(res)
Результат выполнения кода:
True
Пример проверки с использованием разных допусков
Можно регулировать точность проверки через параметры atol
(абсолютная погрешность) и rtol (относительная погрешность):
import torch
from torch.autograd import gradgradcheck
def func(x):
return torch.sin(x) * torch.cos(x)
torch.manual_seed(3)
x = torch.randn(3, requires_grad=True)
res = gradgradcheck(
func,
(x,),
atol=1e-4,
rtol=1e-2
)
print(res)
Результат выполнения кода:
True
Смотрите также
-
функцию
gradcheck,
которая проверяет только первые производные пользовательской функции -
функцию
backward,
которая вычисляет градиенты для выходных тензоров -
функцию
grad,
которая вычисляет градиенты для указанных входных тензоров -
контекстный менеджер
detect_anomaly,
который помогает обнаруживать аномалии в вычислении градиентов