РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
71 of 769 menu

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,
    который показывает, сохраняется ли градиент для нелистовых тензоров
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить