Метод register_forward_pre_hook
Метод register_forward_pre_hook класса Module
регистрирует функцию-хук, которая будет вызываться непосредственно
перед выполнением forward-прохода модуля. Хук принимает три аргумента:
сам модуль, кортеж входных аргументов и кортеж входных аргументов
в формате ключевых слов. Зарегистрированные хуки полезны для
отладки, логирования, визуализации или модификации входных данных
перед их передачей в модуль.
Синтаксис
module.register_forward_pre_hook(hook)
Где hook - функция, которая будет вызвана перед forward-проходом.
Метод возвращает объект-дескриптор, который можно использовать для
удаления зарегистрированного хука.
Пример
Давайте создадим простой линейный слой и зарегистрируем хук, который выводит размерности входных данных перед forward-проходом:
import torch
import torch.nn as nn
def pre_hook(module, input_args, input_kwargs):
print(f"Input shape: {input_args[0].shape}")
linear = nn.Linear(10, 5)
hook_handle = linear.register_forward_pre_hook(pre_hook)
t = torch.randn(3, 10)
res = linear(t)
hook_handle.remove()
Результат выполнения кода:
"Input shape: torch.Size([3, 10])"
Хук успешно перехватил входные данные перед вычислениями и вывел их размерность.
Пример
Давайте используем хук для модификации входных данных перед их обработкой модулем. В этом примере мы добавим случайный шум к входным данным:
import torch
import torch.nn as nn
torch.manual_seed(0)
def noise_hook(module, input_args, input_kwargs):
x = input_args[0]
noise = torch.randn_like(x)
x_noisy = x + noise
return (x_noisy,)
linear = nn.Linear(10, 5)
hook_handle = linear.register_forward_pre_hook(noise_hook)
t = torch.ones(2, 10)
res = linear(t)
print(res)
hook_handle.remove()
Результат выполнения кода:
tensor([
[1.2994, 0.2998, 0.0899, 1.5260, 0.1059],
[1.2994, 0.2998, 0.0899, 1.5260, 0.1059],
], grad_fn=<AddmmBackward0>)
Хук модифицировал входные данные, добавив к ним шум, и вернул новый кортеж, который был передан в модуль вместо исходных данных.
Пример
Рассмотрим более сложный пример, где мы применяем хук для логирования вложенных модулей. Создадим последовательную модель и зарегистрируем предварительный хук для каждого слоя:
import torch
import torch.nn as nn
def logging_hook(module, input_args, input_kwargs):
module_name = type(module).__name__
x = input_args[0]
print(f"{module_name}: input norm = {x.norm().item():.4f}")
model = nn.Sequential(
nn.Linear(10, 20),
nn.ReLU(),
nn.Linear(20, 5)
)
handles = []
for child in model:
handle = child.register_forward_pre_hook(logging_hook)
handles.append(handle)
t = torch.randn(4, 10)
res = model(t)
for handle in handles:
handle.remove()
Результат выполнения кода:
"Linear: input norm = 3.9766"
"ReLU: input norm = 5.4547"
"Linear: input norm = 4.5233"
Для каждого слоя был вызван хук, который вывел норму входных данных, что позволяет отслеживать изменения данных по мере прохождения через модель.
Смотрите также
-
метод
register_forward_hook,
который регистрирует хук, вызываемый после forward-прохода -
метод
forward,
который определяет прямое распространение модуля -
метод
apply,
который рекурсивно применяет функцию ко всем вложенным модулям -
метод
register_full_backward_hook,
который регистрирует хук для обратного распространения ошибки