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

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