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

Метод register_full_backward_hook

Метод register_full_backward_hook класса Module регистрирует хук, который будет вызван после того, как градиенты для всех входов и выходов модуля будут полностью вычислены. В отличие от обычного register_backward_hook, данный метод получает полные градиенты для всех входов и выходов, включая градиенты для не-тензорных аргументов. Первым аргументом метод принимает функцию-обработчик, которая будет вызываться с тремя параметрами: модуль, градиенты на входе и градиенты на выходе. Метод возвращает дескриптор, который можно использовать для удаления хука.

Синтаксис

hook_handle = module.register_full_backward_hook(hook_function)

Параметр hook_function представляет собой функцию с сигнатурой:

def hook(module, grad_input, grad_output): # module - экземпляр модуля # grad_input - кортеж градиентов по входам # grad_output - кортеж градиентов по выходам pass

Функция-обработчик не должна изменять переданные градиенты, так как это может нарушить процесс обратного распространения. Однако она может использовать их для отладки, логирования или анализа.

Пример

Давайте создадим простой линейный слой и зарегистрируем хук для отслеживания градиентов во время обратного распространения:

import torch import torch.nn as nn class SimpleModel(nn.Module): def __init__(self): super().__init__() self.fc = nn.Linear(5, 3) def forward(self, x): return self.fc(x) def backward_hook(module, grad_input, grad_output): print("Module:", module) print("Grad input:", grad_input) print("Grad output:", grad_output) model = SimpleModel() handle = model.register_full_backward_hook(backward_hook) x = torch.randn(2, 5, requires_grad=True) y = model(x) loss = y.sum() loss.backward() handle.remove()

Результат выполнения кода:

Module: SimpleModel( (fc): Linear(in_features=5, out_features=3, bias=True) ) Grad input: (None,) Grad output: (tensor([[1., 1., 1.], [1., 1., 1.]]),)

В данном примере хук вызывается после вычисления градиентов для слоя. grad_input содержит градиенты по входу модуля, а grad_output - градиенты по выходу.

Пример

Используем хук для анализа градиентов в более сложной модели с несколькими слоями:

import torch import torch.nn as nn class MultiLayerModel(nn.Module): def __init__(self): super().__init__() self.fc1 = nn.Linear(10, 5) self.fc2 = nn.Linear(5, 3) def forward(self, x): x = torch.relu(self.fc1(x)) return self.fc2(x) def hook_fn(module, grad_input, grad_output): print(f"Hook for {module.__class__.__name__}") print(f"Grad input shapes: {[g.shape if g is not None else None for g in grad_input]}") print(f"Grad output shapes: {[g.shape if g is not None else None for g in grad_output]}\n") model = MultiLayerModel() handle1 = model.fc1.register_full_backward_hook(hook_fn) handle2 = model.fc2.register_full_backward_hook(hook_fn) x = torch.randn(4, 10, requires_grad=True) y = model(x) loss = y.mean() loss.backward() handle1.remove() handle2.remove()

Результат выполнения кода:

Hook for Linear Grad input shapes: [None, torch.Size([10, 5])] Grad output shapes: [torch.Size([4, 3])] Hook for Linear Grad input shapes: [None, torch.Size([5, 3])] Grad output shapes: [torch.Size([4, 5])]

Обратите внимание, что grad_input может содержать None для тех входов, которые не требуют градиентов (например, первый элемент кортежа соответствует градиенту по входному тензору, который в данном случае требует градиентов, но в других случаях может быть None).

Пример

Продемонстрируем удаление зарегистрированного хука с помощью возвращаемого дескриптора:

import torch import torch.nn as nn model = nn.Linear(4, 2) def hook_fn(module, grad_input, grad_output): print("Hook is called!") handle = model.register_full_backward_hook(hook_fn) x = torch.randn(3, 4, requires_grad=True) y = model(x) loss = y.sum() loss.backward() # Hook будет вызван handle.remove() # Удаляем хук x2 = torch.randn(3, 4, requires_grad=True) y2 = model(x2) loss2 = y2.sum() loss2.backward() # Hook не будет вызван print("Backward completed without hook")

Результат выполнения кода:

Hook is called! Backward completed without hook

Как видно из примера, после вызова handle.remove хук перестает вызываться при обратном распространении.

Смотрите также

  • метод register_forward_hook,
    который регистрирует хук, вызываемый во время прямого прохода
  • метод register_forward_pre_hook,
    который регистрирует хук, вызываемый перед прямым проходом
  • метод zero_grad,
    который обнуляет градиенты всех параметров модуля
  • метод apply,
    который рекурсивно применяет функцию ко всем подмодулям
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить