Функция gradcheck
Функция gradcheck проверяет правильность вычисления градиентов
пользовательской функции путём сравнения аналитических и численных градиентов.
Она необходима для отладки пользовательских функций и слоёв,
использующих автоградирование. Первым параметром функция принимает
проверяемую функцию, вторым - кортеж входных тензоров,
для которых требуется проверить градиенты. Функция возвращает
логическое значение True, если градиенты вычислены верно,
иначе выбрасывает исключение.
Синтаксис
torch.autograd.gradcheck(
func,
inputs,
eps=1e-6,
atol=1e-5,
rtol=1e-3,
nondet_tol=0.0,
check_undefined_grad=True,
check_grad_dtypes=True,
check_batched_grad=False,
raise_exception=True,
)
Основные параметры
Функция принимает следующие ключевые параметры:
func - проверяемая функция, которая принимает входные тензоры
и возвращает тензор или кортеж тензоров.
inputs - кортеж входных тензоров, для которых вычисляются градиенты.
eps - размер шага для численного дифференцирования (по умолчанию 1e-6).
atol - абсолютная допустимая погрешность (по умолчанию 1e-5).
rtol - относительная допустимая погрешность (по умолчанию 1e-3).
Пример с простой функцией
Проверим градиенты функции, вычисляющей сумму квадратов элементов:
import torch
from torch.autograd import gradcheck
def sum_of_squares(x):
return (x ** 2).sum()
t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
res = gradcheck(sum_of_squares, (t,))
print(res)
Результат выполнения кода:
True
Пример с пользовательским слоем
Проверим градиенты пользовательского линейного слоя:
import torch
from torch.autograd import gradcheck
class CustomLinear(torch.nn.Module):
def __init__(self, in_features, out_features):
super().__init__()
self.weight = torch.nn.Parameter(
torch.randn(out_features, in_features)
)
self.bias = torch.nn.Parameter(
torch.randn(out_features)
)
def forward(self, x):
return torch.mm(x, self.weight.t()) + self.bias
torch.manual_seed(0)
layer = CustomLinear(3, 2)
x = torch.randn(2, 3, requires_grad=True)
def test_func(x):
return layer(x)
res = gradcheck(test_func, (x,))
print(res)
Результат выполнения кода:
True
Пример с функцией нескольких переменных
Проверим градиенты функции, зависящей от двух тензоров:
import torch
from torch.autograd import gradcheck
def complex_func(x, y):
return (x ** 3).sum() + (y ** 2).sum()
t1 = torch.tensor([1.0, 2.0], requires_grad=True)
t2 = torch.tensor([3.0, 4.0], requires_grad=True)
res = gradcheck(complex_func, (t1, t2))
print(res)
Результат выполнения кода:
True
Пример с использованием gradgradcheck
Проверим корректность вычисления градиентов второго порядка
с помощью функции gradgradcheck:
import torch
from torch.autograd import gradgradcheck
def quadratic_func(x):
return (x ** 2).sum()
t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
res = gradgradcheck(quadratic_func, (t,))
print(res)
Результат выполнения кода:
True
Пример с изменением точности проверки
Настроим параметры точности для проверки градиентов:
import torch
from torch.autograd import gradcheck
def power_func(x):
return (x ** 4).sum()
t = torch.tensor([0.5, 1.5, 2.5], requires_grad=True)
res = gradcheck(
power_func,
(t,),
eps=1e-5,
atol=1e-4,
rtol=1e-2
)
print(res)
Результат выполнения кода:
True
Смотрите также
-
функцию
backward,
которая вычисляет градиенты для тензоров -
функцию
grad,
которая вычисляет градиенты выходов по входам -
функцию
gradgradcheck,
которая проверяет градиенты второго порядка -
контекстный менеджер
no_grad,
который отключает вычисление градиентов