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.
- 8 min read
- 3 reading levels
- Published
Read these first
On this page 5
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 learnEach 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
- Writing a custom loss function — shaping the gradients you now know how to watch.
- Backpropagation — the river the backward hooks stand beside.
- Weight initialisation — the usual culprit when hook readings look wrong at step zero.
Developer — Code and libraries.
Setup
pip install torchWritten 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
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))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
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()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
- Writing a custom loss function — shaping the gradients you now know how to watch.
- Backpropagation — the river the backward hooks stand beside.
- Weight initialisation — the usual culprit when hook readings look wrong at step zero.
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_extractorsupersedes 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
- Writing a custom loss function — shaping the gradients you now know how to watch.
- Backpropagation — the river the backward hooks stand beside.
- Weight initialisation — the usual culprit when hook readings look wrong at step zero.