Building Models with nn.Module

Hooks for inspecting a running model

Hooks are small functions PyTorch calls for you as data flows through each layer — a way to watch activations and gradients anywhere in a model without editing its code.

On this page 5
  1. Why it exists
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

One lesson, three depths. Pick the one that fits you today — you can switch any time.

Beginner — No maths. Plain English.

A hook is a function you attach to a layer, and PyTorch runs it automatically every time data passes through that layer.

Think of CCTV cameras in a factory. The production line runs exactly as before — nothing slows down, no station is rebuilt. But now you can watch station 7 from the control room and see what actually passes through it.

Hooks are those cameras for a neural network. You clip one onto a layer, and every time the layer does its work, your function gets shown the input and the output. The model's own code is never touched.

Why it exists

A deep model is a black box mid-flight. Data goes in, an answer comes out, and the dozens of layers in between are silent. When something is wrong — every prediction identical, training stuck — the question is always where inside.

You could edit the model to add print statements. But often you cannot: the model came from a library, or is 200 layers deep, or the edit itself might change behaviour. Hooks solve this politely. Attach a camera, watch, unclip it when done.

They watch in both directions. Forward hooks see the data flowing toward the answer. Backward hooks see the learning signal flowing back — the corrections described in backpropagation.

How it works

            forward direction →
  input → [layer 1] → [layer 2] → [layer 3] → output
              |            |           |
           (camera)     (camera)    (camera)     ← your hooks
              |            |           |
        "shapes ok"  "half are zero"  "all zero!"   ← what you learn

Each camera reports without interfering. Finding the layer where signals die is exactly how real debugging sessions go.

A real example you have seen

Heat-map explanations of AI decisions — the images showing where a medical model looked when it flagged a scan — are made with hooks. A camera on a late layer records what activated, and that recording becomes the highlighted picture.

Remember this

  • A hook is a function PyTorch calls for you whenever a layer runs.
  • Forward hooks watch data; backward hooks watch the learning signal.
  • The model's code is never edited — attach, observe, remove.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install torch

Written and tested against torch 2.5 on CPU. Exact percentages below depend on the seed and version; the shapes do not.

A forward hook hunting dead ReLUs

forward_hooks.py
import torch
from torch import nn

torch.manual_seed(0)
model = nn.Sequential(
    nn.Linear(16, 32), nn.ReLU(),
    nn.Linear(32, 32), nn.ReLU(),
    nn.Linear(32, 2),
)

captured = {}

def watch(name):
    def hook(module, inputs, output):
        # runs after this module's forward; we look, we do not touch
        captured[name] = (tuple(output.shape), (output == 0).float().mean().item())
    return hook

handles = [layer.register_forward_hook(watch(f"layer{i}"))
           for i, layer in enumerate(model) if isinstance(layer, nn.ReLU)]

model(torch.randn(64, 16))
for name, (shape, dead) in captured.items():
    print(f"{name}: output {shape}, {dead:.0%} of activations are zero")

for h in handles:
    h.remove()                       # a forgotten hook keeps running forever
print("hooks removed:", all(len(m._forward_hooks) == 0 for m in model))
Output
layer1: output (64, 32), 51% of activations are zero
layer3: output (64, 32), 54% of activations are zero
hooks removed: True

Around half the activations are zero — healthy for ReLU on random data. In a sick network you will see 95% and rising, which is the "dying ReLU" disease from activation functions, caught red-handed in two dozen lines.

The anatomy

The signature is fixed: hook(module, inputs, output). inputs is a tuple (layers can take several arguments); output is whatever forward returned. Return None to observe. Returning a tensor replaces the output — powerful, and a footgun if you did it by accident.

The closure over name is the standard trick for telling cameras apart. register_forward_hook accepts one function; wrapping it in watch(name) bakes the label in.

The handle returned by every register_* call has one method, remove(). Keep the handles. A hook left attached runs on every forward pass forever — in training, that is millions of surprise function calls, and if the hook stores tensors, a memory leak.

Backward hooks: watching the gradient river

backward_hooks.py
import torch
from torch import nn

torch.manual_seed(0)
model = nn.Sequential(nn.Linear(16, 32), nn.Tanh(), nn.Linear(32, 1))

def grad_hook(name):
    def hook(module, grad_input, grad_output):
        print(f"{name:12s} mean |grad arriving|: {grad_output[0].abs().mean().item():.4f}")
    return hook

for i, layer in enumerate(model):
    layer.register_full_backward_hook(grad_hook(f"{i}:{type(layer).__name__}"))

loss = model(torch.randn(8, 16)).sum()
loss.backward()
Output
2:Linear     mean |grad arriving|: 1.0000
1:Tanh       mean |grad arriving|: 0.0936
0:Linear     mean |grad arriving|: 0.0690

The printout runs backwards — last layer first — because that is the direction gradients flow. Watch the size shrink as it passes through Tanh. Stack twenty such layers and this shrinkage becomes the vanishing-gradient story told with live numbers.

Use register_full_backward_hook, not the older register_backward_hook — the old one has documented incorrect behaviour on multi-input modules and survives only for compatibility.

Common mistakes

Storing outputs with their graphs. captured[name] = output keeps the whole autograd graph alive; memory climbs every batch. Store output.detach() (or better, computed statistics) unless you need the graph.

Forgetting remove(). Especially in notebooks, where re-running a cell attaches a second hook alongside the first. Symptoms: duplicated prints, slowing epochs.

Mutating in a hook by accident. In-place edits to output inside a forward hook change what downstream layers receive. Observe with .detach(), modify only on purpose.

Hooking the wrong granularity. A hook on model fires once for the whole model. Hook the specific submodule you care about — model[2], or model.encoder.layer4.

Try it yourself

Attach forward hooks to the Linear layers instead, and record output.abs().mean(). Then re-initialise the model with weights scaled 20x larger and watch the averages explode layer by layer — weight initialisation demonstrated by camera.

What to learn next

Researcher — Mathematics and papers.

The full hook surface

Module-level, in execution order: register_forward_pre_hook (before forward; may rewrite inputs), register_forward_hook (after; may rewrite output), register_full_backward_pre_hook (before the module's backward), register_full_backward_hook (after; sees grad_input, grad_output). Since torch 2.0, register_state_dict_pre_hook and friends extend the same pattern to serialization, and nn.modules.module.register_module_forward_hook installs a global hook on every module — how profilers instrument models they have never seen.

Tensor-level: tensor.register_hook(fn) fires when that tensor's gradient is computed, receiving and optionally replacing grad. This is the primitive underneath gradient clipping experiments, per-tensor gradient logging, and gradient reversal layers (Ganin and Lempitsky, 2015, domain-adversarial training) — though the cleaner implementation of gradient reversal is a custom autograd.Function.

Semantics worth precision: grad_output is the gradient of the loss w.r.t. the module's outputs; grad_input w.r.t. its inputs — for a module $f$ with $y = f(x)$, the hook observes $\partial L/\partial y$ and $\partial L/\partial x$, i.e. both ends of the module's slice of the chain rule. The legacy register_backward_hook mis-reported grad_input for modules with multiple inputs, hence the full replacement (torch 1.8+).

What hooks built

  • Grad-CAM (Selvaraju et al., 2017): forward hook stores the last conv feature map; backward hook stores its gradient; their weighted combination localises the evidence for a class.
  • Feature extraction: torchvision.models.feature_extraction.create_feature_extractor supersedes hand-rolled hooks via FX graph rewriting, but hooks remain the fallback for dynamic models FX cannot trace.
  • Activation statistics at scale: outlier-feature studies in LLMs (Dettmers et al., 2022, LLM.int8()) rest on exactly the hook pattern above, applied to transformer blocks.
  • Mechanistic interpretability: activation patching and steering (e.g. TransformerLens) is hooks as a research method — intervene on output, observe downstream behaviour.

Interaction with the compiled world

Hooks are Python callbacks woven into eager execution, and they constrain acceleration: torch.compile supports module hooks partially (graph breaks around them, with coverage improving by version — verify against your version's docs before relying on it), and anything that inspects .grad mid-backward serialises against autograd's execution. The profiling-grade alternative is torch.profiler with record_function ranges; hooks are for semantic observation, profilers for timing.

Reading

  • Selvaraju et al. (2017), Grad-CAM: Visual Explanations from Deep Networks via Gradient-based Localization.
  • Ganin and Lempitsky (2015), Unsupervised Domain Adaptation by Backpropagation.
  • PyTorch docs, "Module Hooks" — the normative statement of signatures and ordering.

What to learn next