torch.autograd.gradgradcheck
-
torch.autograd.gradgradcheck(func, inputs, grad_outputs=None, *, eps=1e-06, atol=1e-05, rtol=0.001, gen_non_contig_grad_outputs=False, raise_exception=True, nondet_tol=0.0, check_undefined_grad=True, check_grad_dtypes=False, check_batched_grad=False, check_fwd_over_rev=False, check_rev_over_rev=True, fast_mode=False, masked=False)[source] -
Проверка градиентов градиентов, вычисленных с помощью малых конечных разностей, по аналитическим градиентам относительно тензоров в
inputsиgrad_outputs, которые имеют тип с плавающей точкой или комплексным типом и сrequires_grad=True.Эта функция проверяет, правильна ли обратная проходка через вычисленные градиенты для заданных
grad_outputs.Проверка между численными и аналитическими градиентами использует
allclose().Примечание
Значения по умолчанию предназначены для
inputиgrad_outputsдвойной точности. Эта проверка, вероятно, потерпит неудачу, если они имеют меньшую точность, например,FloatTensor.Предупреждение
Если какой-либо проверенный тензор в
inputиgrad_outputsимеет перекрывающуюся память, т.е. разные индексы, указывающие на один и тот же адрес памяти (например, изtorch.expand()), эта проверка, вероятно, потерпит неудачу, потому что численные градиенты, вычисленные точечной пертурбацией в таких индексах, изменят значения во всех других индексах, которые используют тот же адрес памяти.- Параметры
-
- func (функция) – Python-функция, которая принимает тензорные входные данные и возвращает тензор или кортеж тензоров
- inputs (кортеж из Tensor или Tensor) – входные данные для функции
- grad_outputs (кортеж из Tensor или Tensor, необязательно) – Градиенты по отношению к выходным данным функции.
- eps (число с плавающей точкой, необязательно) – возмущение для конечных разностей
- atol (число с плавающей точкой, необязательно) – абсолюльная погрешность
- rtol (число с плавающей точкой, необязательно) – относительная погрешность
-
gen_non_contig_grad_outputs (булево значение, необязательно) – если
grad_outputsравноNoneиgen_non_contig_grad_outputsравноTrue, случайно сгенерированные выходные градиенты делают несмежными - raise_exception (булево значение, необязательно) – указывает, следует ли генерировать исключение в случае неудачи проверки. Исключение содержит более подробную информацию о характере ошибки. Это полезно при отладке gradchecks.
- nondet_tol (число с плавающей точкой, необязательно) – погрешность для недетерминизма. При запуске одинаковых входных данных через дифференцирование результаты должны либо совпадать точно (по умолчанию, 0,0), либо находиться в пределах этой погрешности. Обратите внимание, что небольшое количество недетерминизма в градиенте приведет к более крупным погрешностям во второй производной.
- check_undefined_grad (булево значение, необязательно) – если True, проверить, поддерживаются ли неопределенные выходные градиенты и обрабатываются как нули
- check_batched_grad (булево значение, необязательно) – если True, проверить, можем ли мы вычислить пакетные градиенты с помощью поддержки vmap прототипов. По умолчанию False.
- fast_mode (булево значение, необязательно) – если True, запустить более быструю реализацию gradgradcheck, которая больше не вычисляет весь якобиан.
- masked (булево значение, необязательно) – если True, градиенты неопределённых элементов разреженных тензоров игнорируются (по умолчанию, False).
- Возвращает
-
True, если все различия удовлетворяют условию allclose
- Тип возвращаемого значения
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.autograd.gradgradcheck.html