How do you capture a PyTorch model's intermediate activations without editing its forward()?
answer
- callback, not a source edit
- fires from __call__, after forward returns
- (module, args, output) signature
- detach or you pin the graph
- the handle exists to be removed
basics
~20 sRegister a forward hook on the submodule you want: handle = layer.register_forward_hook(fn), where fn receives (module, args, output). Detach what you store, and call handle.remove() afterwards or the hook keeps firing and holding tensors alive.
solid answer
~40 s`nn.Module.register_forward_hook(fn)` attaches a callback that `__call__` invokes after `forward()` returns, with the signature `(module, args, output)`. That lets you record — or even replace, by returning a new value — any submodule's output without touching the model's source, which matters when the model comes from a library you do not own. Reach the submodule with `model.get_submodule("encoder.layers.3")` or by walking `named_modules()`. Two production cautions: store `output.detach()`, otherwise you pin the whole autograd graph for that step and memory climbs every iteration; and keep the returned `RemovableHandle` so you can call `handle.remove()`, since a forgotten hook keeps firing in later phases and accumulating tensors. `register_forward_pre_hook` does the same on the way in, and hooks fire only through `model(x)` — a direct `.forward()` call skips them entirely.
code
python · 18 linesimport torch
import torch.nn as nn
model = nn.Sequential(nn.Linear(8, 16), nn.ReLU(), nn.Linear(16, 4))
acts = {}
def save(name):
def hook(module, args, output):
acts[name] = output.detach()
return hook
handle = model[0].register_forward_hook(save("fc1"))
try:
model(torch.randn(2, 8))
finally:
handle.remove()
print(acts["fc1"].shape) # torch.Size([2, 16])go deeper
Know that register_forward_hook attaches a callback receiving (module, args, output), and that it returns a handle you are expected to remove when done.
Explain where hooks fit — dispatched by nn.Module.call around forward — and how to reach a nested layer by dotted path with get_submodule or by walking named_modules.
Demonstrate the operational discipline: detach what you store or you retain the graph, remove handles in a finally block, and recognise a dead hook as either a direct forward() call or the wrong module instance.
Decide when hooks are the wrong tool at all: for a permanent extraction path shipped to production, express intermediates in forward or use a graph-level rewrite, since hooks are Python callbacks that no exported artifact carries.
## The problem hooks solve You need the activations of layer 3 of a pretrained encoder — for a probing experiment, a feature extractor, a similarity search index, or just to see where a signal dies. The obvious approaches are bad: forking the library to add a `return_intermediates=True` flag couples you to that fork forever, and rebuilding the model as a `nn.Sequential` of its children silently drops anything the original `forward` did between modules (reshapes, residual adds, functional activations). Hooks are PyTorch's supported answer. They are callbacks the module dispatches during `__call__`, so they see real traffic through the real `forward`. ## The API `handle = module.register_forward_hook(fn)` where `fn(module, args, output)`: - `module` is the module the hook is attached to; - `args` is the tuple of positional inputs it was called with; - `output` is whatever `forward` returned. Returning `None` (the usual case) leaves the output untouched; returning a value **replaces** the module's output for the rest of the computation, which is how people patch activations for interpretability experiments. `register_forward_pre_hook(fn)` runs before `forward`, with `fn(module, args)`, and can replace the inputs. `register_full_backward_hook` exists for the gradient side. Both registration calls return a `RemovableHandle` whose `.remove()` deregisters the hook. That handle is the piece people forget. ## Finding the submodule Two reliable routes: - `model.get_submodule("encoder.layers.3.mlp")` — dotted-path lookup, raises `AttributeError` on a bad path, which is exactly what you want rather than silently hooking nothing. - iterate `model.named_modules()` and match on the name or on `isinstance`. This is how you hook *every* `nn.Conv2d` in one loop for a layer-wise statistics dump. The names are the same attribute paths that form `state_dict` keys, so anything you can see in a checkpoint you can address here. ## The two things that go wrong in production **Holding the graph.** If your hook does `store[name] = output`, you keep a tensor that still carries `grad_fn`, and through it the entire autograd graph for that iteration. Backward cannot free what you are still referencing. Memory climbs every step until the process OOMs — and because the growth is gradual it reads exactly like a slow leak elsewhere. Store `output.detach()`, and `.cpu()` too if you are collecting across many steps. **Leaked handles.** Hooks registered in a helper and never removed keep firing through validation, through export tracing, through the next experiment in the same process. A hook that appends to a list becomes an unbounded accumulator. The disciplined shape is a try/finally, or a small context manager that removes every handle it registered on exit. ## The silent no-op Hooks fire from `nn.Module.__call__`. If any code path invokes `submodule.forward(x)` directly, the hook does not run. You get an empty dictionary and no error at all. The same applies if you attached the hook to the wrong instance — for example to a freshly constructed layer rather than the one inside the model, or to `ddp_model.module`'s child while calling a differently wrapped object. When a hook seems dead, verify the identity of the module you hooked and confirm the call path goes through `model(x)`. ## Alternatives and when they are better - **`torch.fx` / `create_feature_extractor`-style graph rewriting** gives you a new module whose forward returns the nodes you asked for. It is more robust for a fixed extraction you will run many times, and it survives serialization; it is heavier to set up and can fail on models with data-dependent control flow. - **Subclassing and overriding `forward`** is fine when you own the model. Hooks earn their keep precisely when you do not. - **Export-time capture** is a different problem: a hook is a Python callback, so it does not survive `torch.export` or a traced graph. If your goal is to ship a model that returns intermediates, put that in the model, not in a hook. ## The shape of good hook code Register, run, remove, in a bounded scope. Detach what you keep. Address submodules by their dotted path so the code fails loudly when the architecture changes underneath you. Never leave a hook installed across phases of a job — activations captured during validation because a training hook was never removed is a genuinely painful bug to find.
- Why does storing the raw output in a forward hook cause memory to climb?The output tensor still carries `grad_fn`, so a reference to it keeps the whole autograd graph for that iteration alive. Backward frees the graph only if nothing else holds it. Every step adds another retained graph and memory grows steadily until an OOM that looks like a leak somewhere else. `output.detach()` — plus `.cpu()` when collecting across steps — fixes it.
- What happens if a forward hook returns a value instead of None?The returned value replaces the module's output for the rest of the computation. That is a supported feature, used for activation patching and ablation experiments — zero a head's output, or substitute an activation from a different input, without editing the model. It is also a trap if you accidentally return something from a hook meant only to record.
- Your hook never fires and no error is raised — what do you check?First, whether the call path goes through `model(x)`; hooks dispatch from `__call__`, so a direct `.forward()` call anywhere in the chain skips them. Second, module identity: confirm you hooked the instance inside the model, ideally via `model.get_submodule("path.to.layer")` rather than a separately constructed layer or a differently wrapped object.
- Do forward hooks survive torch.export or a TorchScript trace?No. A hook is arbitrary Python invoked at call time, not part of the module's computation graph, so exported or traced artifacts do not carry it. If the deployed model must return intermediates, that has to be expressed in `forward` itself; hooks are a development and analysis tool, not a serialization feature.
saying these in an interview costs you the question
- Stores the raw output and pins the autograd graph
- Never calls handle.remove(), so hooks leak across phases
- Expects hooks to fire on a direct forward() call
- Rebuilds the model as Sequential and loses functional ops
- Assumes hooks survive export or tracing