register_post_accumulate_grad_hook
Метод register_post_accumulate_grad_hook регистрирует функцию-хук,
которая будет вызвана автоматически после того, как градиент для тензора
был накоплен в его атрибуте grad. Хук получает на вход единственный
аргумент - тензор градиента, накопленный в grad. Этот метод полезен
для модификации градиента перед его применением оптимизатором, например,
для клиппинга, нормализации или применения специальных преобразований.
Важно: хук срабатывает только для листовых тензоров (тензоров с
requires_grad=True, которые являются листьями графа вычислений),
так как только для них накапливается градиент. Возвращаемое значение
хука игнорируется, изменения применяются непосредственно к атрибуту
grad внутри хука.
Синтаксис
tensor.register_post_accumulate_grad_hook(hook_fn)
где hook_fn - функция, принимающая один аргумент (тензор градиента)
и не возвращающая ничего (или возвращающая None).
Пример
Простейший пример регистрации хука, который выводит информацию о градиенте:
import torch
def print_grad_hook(grad):
print(f"Gradient: {grad}")
t = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
hook_handle = t.register_post_accumulate_grad_hook(print_grad_hook)
loss = (t ** 2).sum()
loss.backward()
# Хук будет вызван после накопления градиента
print(t.grad)
Результат выполнения кода:
Gradient: tensor([2., 4., 6.])
tensor([2., 4., 6.])
Пример
Хук может модифицировать градиент на месте, например, применять клиппинг:
import torch
def clip_grad_hook(grad):
# Ограничиваем значения градиента диапазоном [-2, 2]
grad.clamp_(-2, 2)
t = torch.tensor([-5.0, 0.0, 5.0], requires_grad=True)
hook_handle = t.register_post_accumulate_grad_hook(clip_grad_hook)
loss = (t ** 2).sum()
loss.backward()
print(t.grad)
Результат выполнения кода:
tensor([-2., 0., 2.])
Пример
Хуки полезны для отладки или анализа градиентов в процессе обучения. В этом примере хук накапливает статистику градиентов:
import torch
def grad_stats_hook(grad):
grad_mean = grad.mean().item()
grad_std = grad.std().item()
print(f"Grad stats - mean: {grad_mean:.4f}, std: {grad_std:.4f}")
# Создаём тензор с случайными значениями
torch.manual_seed(0)
t = torch.randn(10, requires_grad=True)
hook_handle = t.register_post_accumulate_grad_hook(grad_stats_hook)
# Вычисляем и обратно распространяем потерю
loss = (t ** 3).sum()
loss.backward()
Результат выполнения кода:
Grad stats - mean: 2.1757, std: 2.3448
Пример
Хук можно удалить, вызвав метод remove у объекта-обработчика:
import torch
def my_hook(grad):
print("Hook called!")
t = torch.tensor([1.0, 2.0], requires_grad=True)
handle = t.register_post_accumulate_grad_hook(my_hook)
# Первый backward вызовет хук
loss1 = (t ** 2).sum()
loss1.backward()
# Удаляем хук
handle.remove()
# Второй backward не вызовет хук
t.grad.zero_() # обнуляем градиент для повторного расчёта
loss2 = (t ** 3).sum()
loss2.backward()
Результат выполнения кода:
Hook called!
Смотрите также
-
метод
register_hook,
который регистрирует хук, вызываемый во время обратного распространения -
метод
backward,
который запускает вычисление градиентов -
атрибут
grad,
в котором хранится накопленный градиент тензора -
атрибут
retains_grad,
который показывает, сохраняется ли градиент для нелистовых тензоров