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

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