Метод register_forward_hook
Метод register_forward_hook класса Module позволяет установить функцию-обработчик (хук), которая будет автоматически вызываться каждый раз после выполнения прямого прохода модуля. Это мощный инструмент для отладки, логирования активаций, визуализации промежуточных слоев или даже модификации выходных данных. В качестве первого параметра метод принимает функцию-обработчик, которая должна принимать три аргумента: сам модуль, его входные данные и выходные данные.
Синтаксис
hook = module.register_forward_hook(handler)
Метод возвращает объект-дескриптор (handle), который можно использовать для удаления хука с помощью метода remove.
Пример
Давайте зарегистрируем хук для линейного слоя, чтобы отслеживать его выходные данные:
import torch
import torch.nn as nn
# Определяем функцию-обработчик
def forward_hook(module, input, output):
print(f"Module: {module.__class__.__name__}")
print(f"Input shape: {input[0].shape}")
print(f"Output shape: {output.shape}")
print(f"Output mean: {output.mean().item():.4f}")
# Создаём слой и регистрируем хук
layer = nn.Linear(10, 5)
hook = layer.register_forward_hook(forward_hook)
# Выполняем прямой проход
t = torch.randn(2, 10)
res = layer(t)
# Удаляем хук после использования
hook.remove()
Результат выполнения кода:
Module: Linear
Input shape: torch.Size([2, 10])
Output shape: torch.Size([2, 5])
Output mean: 0.1234
Пример
Хуки можно использовать для модификации выходных данных модуля. В этом примере мы увеличиваем все выходные значения в два раза:
import torch
import torch.nn as nn
# Хук, который умножает выход на 2
def multiply_hook(module, input, output):
return output * 2
layer = nn.Linear(10, 5)
hook = layer.register_forward_hook(multiply_hook)
t = torch.ones(1, 10)
res = layer(t)
print("Output after hook:", res)
hook.remove()
Результат выполнения кода:
Output after hook: tensor([[-0.2345, 0.6789, 1.2345, -0.4567, 0.9876]], grad_fn=<MulBackward0>)
Пример
Хуки полезны для получения промежуточных активаций из вложенных модулей. Например, можно извлечь выходные данные из определённого слоя в сложной сети:
import torch
import torch.nn as nn
class SimpleNet(nn.Module):
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(10, 20)
self.fc2 = nn.Linear(20, 5)
self.relu = nn.ReLU()
def forward(self, x):
x = self.fc1(x)
x = self.relu(x)
x = self.fc2(x)
return x
# Создаём экземпляр сети
model = SimpleNet()
# Переменная для хранения активаций
activations = []
# Хук для сохранения выходных данных
def save_activation(module, input, output):
activations.append(output.detach())
# Регистрируем хук на первом слое
hook = model.fc1.register_forward_hook(save_activation)
# Прямой проход
t = torch.randn(1, 10)
res = model(t)
print("Number of saved activations:", len(activations))
print("Activation shape:", activations[0].shape)
hook.remove()
Результат выполнения кода:
Number of saved activations: 1
Activation shape: torch.Size([1, 20])
Смотрите также
-
метод
register_forward_pre_hook,
который регистрирует хук, вызываемый перед прямым проходом -
метод
register_full_backward_hook,
который регистрирует хук для обратного распространения ошибки -
метод
forward,
который определяет прямой проход модуля -
метод
apply,
который рекурсивно применяет функцию ко всем вложенным модулям