Функция set_detect_anomaly
Функция set_detect_anomaly управляет глобальным режимом обнаружения аномалий
при вычислении градиентов. Когда этот режим включён, PyTorch выполняет дополнительные
проверки в процессе обратного распространения ошибки, что помогает отладить проблемы
с градиентами, такие как NaN или Inf значения, а также ошибки в
вычислительном графе. Первым параметром функция принимает булево значение True
или False, которое включает или отключает режим обнаружения аномалий.
Синтаксис
torch.set_detect_anomaly(mode)
Пример
Давайте включим режим обнаружения аномалий и выполним простое обратное распространение:
import torch
torch.set_detect_anomaly(True)
t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
res = t.sum()
res.backward()
print(t.grad)
Результат выполнения кода:
tensor([1., 1., 1.])
Пример
Режим обнаружения аномалий помогает выявить проблему при возникновении NaN в градиентах:
import torch
torch.set_detect_anomaly(True)
t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
res = t / 0.0
res.backward()
Результат выполнения кода:
RuntimeError: Function 'DivBackward0' returned nan values in its 0th output.
Пример
Давайте отключим режим обнаружения аномалий для оптимизации производительности:
import torch
torch.set_detect_anomaly(False)
t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
res = t.sum()
res.backward()
print(t.grad)
Результат выполнения кода:
tensor([1., 1., 1.])
Смотрите также
-
функцию
detect_anomaly,
которая включает режим обнаружения аномалий в контекстном менеджере -
функцию
backward,
которая вычисляет градиенты для тензора -
функцию
grad,
которая вычисляет градиенты для заданных тензоров -
функцию
gradcheck,
которая проверяет корректность вычисления градиентов