Поиск аномалии градиента в PyTorch
Если обратный проход падает
или градиент обрывается
на неочевидной операции,
помогает режим детекции
аномалии. Функция
set_detect_anomaly
включает или выключает его
для всего процесса.
В активном режиме библиотека подробнее сообщает, где цепочка градиента перестала быть корректной. На обычном графе достаточно обернуть обратный проход:
import torch
x = torch.tensor(2.0, requires_grad=True)
y = x * x
with torch.autograd.set_detect_anomaly(True):
y.backward()
print(x.grad) # выведет tensor(4.)
После отладки режим лучше выключить: он замедляет вычисления. Вызов с ложным флагом снимает контроль:
import torch
torch.autograd.set_detect_anomaly(False)
x = torch.tensor(2.0, requires_grad=True)
y = x * x
y.backward()
print(x.grad) # выведет tensor(4.)
Для 3.0 с записью
постройте квадрат и выполните
обратный проход внутри
включённого режима аномалии.
Выведите градиент аргумента.
Выключите режим аномалии,
затем для 1.0 с записью
удвойте значение и найдите
обратный проход без обёртки.
Выведите поле градиента
аргумента.
Включите режим аномалии,
для двух дробных 2.0
и 0.5 с записью сложите
их и выполните обратный
проход внутри блока.
Выведите градиенты обоих
слагаемых.