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

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