Метод 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,
который рекурсивно применяет функцию ко всем подмодулям