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)[source] -
Проверка градиентов градиентов, вычисленных с помощью малых конечных разностей, по отношению к аналитическим градиентам по тензорам в
inputsиgrad_outputs, которые являются типа с плавающей точкой или комплексным типом и сrequires_grad=True.Эта функция проверяет, что обратное распространение через градиенты, вычисленные для заданных
grad_outputs, является корректным.Проверка между численными и аналитическими градиентами использует
allclose().Примечание
Значения по умолчанию предназначены для
inputиgrad_outputsдвойной точности. Эта проверка, скорее всего, завершится ошибкой, если они имеют меньшую точность, например,FloatTensor.Предупреждение
Если какой-либо проверяемый тензор в
inputиgrad_outputsимеет перекрывающуюся память, т.е. различные индексы указывают на один и тот же адрес памяти (например, изtorch.expand()), эта проверка, скорее всего, завершится ошибкой, потому что численные градиенты, вычисленные путем точечной пертурбации в таких индексах, изменят значения во всех других индексах, которые используют тот же адрес памяти.- Параметры:
-
- func (функция) – функция Python, которая принимает тензорные входные данные и возвращает тензор или кортеж тензоров
- inputs (кортеж тензоров или Tensor) – входные данные функции
- grad_outputs (кортеж тензоров или Tensor, необязательно) – градиенты по отношению к выходным данным функции.
- eps (float, необязательно) – пертурбация для конечных разностей
- atol (float, необязательно) – абсолюльная погрешность
- rtol (float, необязательно) – относительная погрешность
-
gen_non_contig_grad_outputs (bool, необязательно) – если
grad_outputsявляетсяNoneиgen_non_contig_grad_outputsявляетсяTrue, сгенерированные случайным образом градиентные выходы делаются несмежными - raise_exception (bool, необязательно) – указывает, следует ли вызывать исключение при ошибке проверки. Исключение содержит больше информации о точной природе ошибки. Это полезно при отладке gradchecks.
- nondet_tol (float, необязательно) – погрешность для недетерминизма. При выполнении идентичных входных данных через дифференцирование, результаты должны либо совпадать точно (по умолчанию, 0,0), либо быть в пределах этой погрешности. Обратите внимание, что небольшое количество недетерминизма в градиенте приведет к более крупным погрешностям во второй производной.
- check_undefined_grad (bool, необязательно) – если True, проверяется, поддерживаются ли неопределенные градиенты выходов и обрабатываются как нули
- check_batched_grad (bool, необязательно) – если True, проверяется, можем ли мы вычислить пакетные градиенты с помощью прототипа vmap. По умолчанию False.
- fast_mode (bool, необязательно) – если True, выполняется более быстрая реализация gradgradcheck, которая больше не вычисляет всю матрицу Якоби.
- Возвращает:
-
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.gradgradcheck.html