РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
619 of 769 menu

Функция 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,
    который отключает вычисление градиентов
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить