torch.autograd.gradcheck
-
torch.autograd.gradcheck(func, inputs, *, eps=1e-06, atol=1e-05, rtol=0.001, raise_exception=True, check_sparse_nnz=False, nondet_tol=0.0, check_undefined_grad=True, check_grad_dtypes=False, check_batched_grad=False, check_batched_forward_grad=False, check_forward_ad=False, check_backward_ad=True, fast_mode=False)[source] -
Проверка градиентов, вычисленных с помощью малых конечных разностей, по аналитическим градиентам по тензорам в
inputs, которые имеют тип с плавающей точкой или комплексный тип и сrequires_grad=True.Проверка между численными и аналитическими градиентами использует
allclose().Для большинства сложных функций, которые мы рассматриваем для оптимизации, не существует понятия якобиана. Вместо этого gradcheck проверяет, согласованы ли числовые и аналитические значения производных Вирингер и сопряженной Вирингер. Поскольку вычисление градиента выполняется при предположении, что функция в целом имеет вещественный выход, мы обрабатываем функции с комплексным выходом особым образом. Для этих функций gradcheck применяется к двум вещественным функциям, соответствующим взятию вещественных компонент комплексных выходов для первой и взятию мнимых компонент комплексных выходов для второй. Более подробную информацию см. в Автоградиент для комплексных чисел.
Примечание
Значения по умолчанию разработаны для
inputдвойной точности. Эта проверка, вероятно, завершится неудачей, еслиinputимеет меньшую точность, например,FloatTensor.Предупреждение
Если какой-либо проверяемый тензор в
inputимеет перекрывающуюся память, т. е. разные индексы указывают на один и тот же адрес памяти (например, изtorch.expand()), эта проверка, вероятно, завершится неудачей, потому что числовые градиенты, вычисленные точечной пертурбацией в таких индексах, изменят значения во всех других индексах, которые разделяют один и тот же адрес памяти.- Параметры:
-
- func (функция) – функция Python, которая принимает тензорные входные данные и возвращает тензор или кортеж тензоров
- inputs (кортеж из тензора или Tensor) – входные данные функции
- eps (float, необязательно) – пертурбация для конечных разностей
- atol (float, необязательно) – абсолюльная погрешность
- rtol (float, необязательно) – относительная погрешность
- raise_exception (bool, необязательно) – указывает, следует ли поднимать исключение, если проверка завершается неудачей. Исключение предоставляет больше информации об истинном характере ошибки. Это полезно при отладке gradchecks.
- check_sparse_nnz (bool, необязательно) – если True, gradcheck допускает вход SparseTensor, и для любого SparseTensor на входе gradcheck выполнит проверку только в позициях nnz.
- nondet_tol (float, необязательно) – погрешность недетерминированности. При выполнении идентичных входных данных через дифференцирование результаты должны совпадать точно (по умолчанию, 0.0) или находиться в пределах этой погрешности.
-
check_undefined_grad (bool, необязательно) – если True, проверяется, поддерживаются ли неопределенные градиенты выходов и обрабатываются как нули для выходов
Tensor. - check_batched_grad (bool, необязательно) – если True, проверяется, можем ли мы вычислить пакетные градиенты, используя прототип vmap.
- check_batched_forward_grad (bool, необязательно) – если True, проверяет, можем ли мы вычислить пакетные прямые градиенты, используя прямую адресацию и прототип vmap.
- check_forward_ad (bool, необязательно) – если True, проверяется, соответствуют ли градиенты, вычисленные с помощью прямой AD, числовым градиентам.
- check_backward_ad (bool, необязательно) – если False, не выполнять никаких проверок, которые полагаются на обратный режим AD для реализации.
- fast_mode (bool, необязательно) – Быстрый режим для gradcheck и gradgradcheck в настоящее время реализован только для функций R в R. Если ни один из входов и выходов не является комплексным, выполняется более быстрая реализация gradcheck, которая больше не вычисляет весь якобиан; в противном случае мы возвращаемся к медленной реализации.
- Возвращает:
-
True, если все различия удовлетворяют условию allclose
- Тип возвращаемого значения:
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.autograd.gradcheck.html