Функция detect_anomaly
Функция detect_anomaly используется в качестве контекстного менеджера для включения режима обнаружения аномалий в процессе автоматического дифференцирования. В этом режиме PyTorch проверяет градиенты на наличие значений NaN и Inf, и в случае обнаружения выбрасывает ошибку с подробной информацией о месте возникновения. Это мощный инструмент для отладки сложных моделей, когда градиенты неожиданно становятся некорректными.
Синтаксис
torch.autograd.detect_anomaly()
Функция не принимает параметров и возвращает контекстный менеджер.
Пример с аномалией
Рассмотрим код, в котором возникает бесконечное значение во время обратного распространения. Включим режим обнаружения аномалий, чтобы поймать ошибку:
import torch
x = torch.tensor([1.0], requires_grad=True)
y = x / 0 # Деление на ноль приводит к inf
with torch.autograd.detect_anomaly():
y.backward()
Результат выполнения кода:
RuntimeError: Function 'DivBackward0' returned nan values in its 0th output.
Ошибка чётко указывает на функцию, которая произвела некорректное значение.
Пример без ошибок
Если аномалий нет, код выполняется стандартно. Режим detect_anomaly не влияет на корректные вычисления, но добавляет накладные расходы:
import torch
x = torch.tensor([2.0], requires_grad=True)
y = x ** 2
with torch.autograd.detect_anomaly():
y.backward()
print(x.grad)
Результат выполнения кода:
tensor([4.])
Градиент вычислен корректно, ошибок не возникло.
Использование с другими функциями
Режим обнаружения аномалий часто применяется вместе с backward или grad. Это помогает выявить проблемные места в сложных графах вычислений, например, в больших моделях глубокого обучения:
import torch
import torch.nn as nn
class MyModel(nn.Module):
def __init__(self):
super().__init__()
self.linear = nn.Linear(2, 1)
def forward(self, x):
return self.linear(x)
model = MyModel()
x = torch.tensor([[1.0, 2.0]], requires_grad=True)
# Создаем некорректные данные для вызова ошибки
with torch.autograd.detect_anomaly():
output = model(x)
loss = output / 0
loss.backward()
Результат выполнения кода:
RuntimeError: Function 'DivBackward0' returned nan values in its 0th output.
Ошибка указывает на операцию деления, что значительно упрощает поиск бага.
Смотрите также
-
функцию
backward,
которая вычисляет градиенты для тензоров -
функцию
grad,
которая вычисляет градиенты для заданных выходов -
функцию
set_detect_anomaly,
которая включает режим обнаружения аномалий глобально -
функцию
gradcheck,
которая проверяет корректность градиентов через численное дифференцирование