USAAIO Lesson 70. It covers forward and backward hooks for capturing activations and watching gradient flow, memory management with del and gradient checkpointing, torch.profiler for locating training bottlenecks, and the torch.compile modes for faster inference. Every number was verified with torch 2.7.1 and numpy 2.2.6 in June 2026. The lesson runs to 30 slides.
Subject: Machine Learning · 59 slides · code lesson
Open the interactive version of this deck · Homework for this lesson
Title
USAAIO · Lesson 70
Intercept any layer's activations or gradients at runtime, track memory efficiently, find training bottlenecks, and compile models for faster inference.
Objectives
register_forward_hook to capture any layer's output activationsregister_full_backward_hook to inspect and modify gradient flowtorch.utils.checkpoint) to cut activation memory by up to 4xtorch.profiler and tracemalloc to identify the top CPU bottleneck in a training looptorch.compile(model, mode=...) and select the right mode for inference vs trainingWarm-up
Discussion prompt
Before we open Lesson 70: PyTorch Hooks, Profiling & torch.compile: without looking back, what was the main idea of kNN and the Curse of Dimensionality, and what could you do by the end of it that you could not do before?
Hint: One sentence for the idea, one for the skill. If the second one is blank, that is the part to revisit.
Answer:
kNN majority-vote classification from scratch, distance metrics (L2/L1/cosine), the curse of dimensionality (max/min ratio converging to 1), KD-tree and ball-tree for efficient lookup, weighted kNN (1/distance), and bias-variance tradeoff over k. Every number verified with numpy/sklearn, June 2026.
Section
Part 1 of 4
Concept
A forward hook is a callable you attach to any nn.Module. After that module's forward() runs, PyTorch calls your hook with (module, input, output) — live tensors, no re-run needed.
Use cases: visualize feature maps at a CNN layer, compare pre/post-ReLU activations, detect dead neurons, extract embeddings from an intermediate layer without forking the model.
| hook type | signature | typical use |
|---|---|---|
| forward | (module, input, output) | capture or log activations |
| forward pre-hook | (module, input) | modify input before forward |
| full backward | (module, grad_input, grad_output) | inspect / clip gradient flow |
Comparison
Comparison matrix
From What a forward hook is: refill the signature column from what you know. The rest of the table is as it appeared.
| hook type | signature | typical use |
|---|---|---|
| forward | (module, input, output) | capture or log activations |
| forward pre-hook | (module, input) | modify input before forward |
| full backward | (module, grad_input, grad_output) | inspect / clip gradient flow |
Estimation
Predict first
Build a tiny CNN (conv1 -> ReLU -> conv2 -> pool -> fc) and register forward hooks on both conv layers.
Commit before you compute: what does Hooking a CNN to capture activations come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Call h.remove() after each forward pass you need — or use the context manager form
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Uncleaned hooks accumulate: the same hook fires again every forward pass and you collect stale references.
Worked example
Build a tiny CNN (conv1 -> ReLU -> conv2 -> pool -> fc) and register forward hooks on both conv layers.
import torch, torch.nn as nn
class TinyCNN(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(1, 4, 3, padding=1)
self.relu = nn.ReLU()
self.conv2 = nn.Conv2d(4, 8, 3, padding=1)
self.pool = nn.AdaptiveAvgPool2d(1)
self.fc = nn.Linear(8, 10)
def forward(self, x):
x = self.relu(self.conv1(x))
x = self.relu(self.conv2(x))
return self.fc(self.pool(x).flatten(1))
activations = {}
def make_hook(name):
def fn(mod, inp, out): activations[name] = out.detach()
return fn
model = TinyCNN(); model.eval()
h1 = model.conv1.register_forward_hook(make_hook('conv1'))
h2 = model.conv2.register_forward_hook(make_hook('conv2'))
torch.manual_seed(0)
x = torch.randn(1, 1, 8, 8)
with torch.no_grad(): out = model(x)
h1.remove(); h2.remove()
for k, v in activations.items():
print(k, tuple(v.shape), 'mean=%.4f' % v.mean().item())| layer | activation shape | mean (seed=0) |
|---|---|---|
| conv1 | (1, 4, 8, 8) | -0.0590 |
| conv2 | (1, 8, 8, 8) | -0.0013 |
| output logits | (1, 10) | argmax=9 |
Call h.remove() after each forward pass you need — or use the context manager form
Why: Uncleaned hooks accumulate: the same hook fires again every forward pass and you collect stale references. Always remove or scope them.
Trade off
Comparison matrix
From Hooking a CNN to capture activations: every row here is a choice with a cost. Fill the activation shape column, then say which row you would actually pick and what you give up for it.
| layer | activation shape | mean (seed=0) |
|---|---|---|
| conv1 | (1, 4, 8, 8) | -0.0590 |
| conv2 | (1, 8, 8, 8) | -0.0013 |
| output logits | (1, 10) | argmax=9 |
Anomaly
Predict first
A student writes this, and it looks reasonable:
Register register_forward_hook once inside a training loop body to log activations each step.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: Each iteration registers a NEW hook without removing the previous one.
Register once before the loop; remove immediately after sampling.
Why: Each iteration registers a NEW hook without removing the previous one. After 100 epochs you have 100 hooks all firing, storing 100 copies of the activation tensor.
Trap
Register register_forward_hook once inside a training loop body to log activations each step.
for epoch in range(N): h = model.conv1.register_forward_hook(fn); train_step()
Why: Each iteration registers a NEW hook without removing the previous one. After 100 epochs you have 100 hooks all firing, storing 100 copies of the activation tensor.
Register once before the loop; remove immediately after sampling.
h = model.conv1.register_forward_hook(fn); train(); h.remove()
Why: One handle, one hook. PyTorch returns a RemovableHook handle — store it and call .remove() as soon as you are done. Or use a context manager wrapper.
Section
Part 2 of 4
Concept
register_full_backward_hook fires during the backward pass with (module, grad_input, grad_output) — both are tuples of tensors (or None if detached).
Practical uses: gradient flow visualization (are early layers getting gradients?), gradient clipping per-layer, diagnosing vanishing/exploding gradients through a deep network.
Use register_full_backward_hook (not the older register_backward_hook) — the older form is deprecated and misses some autograd nodes when the forward contains multiple branches.
Counterexample
Discussion prompt
register_full_backward_hook fires during the backward pass with (module, grad_input, grad_output) — both are tuples of tensors (or None if detached).
That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.
Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.
Answer:
Practical uses: gradient flow visualization (are early layers getting gradients?), gradient clipping per-layer, diagnosing vanishing/exploding gradients through a deep network.
Missing information
Discussion prompt
Attach full-backward hooks to conv1, conv2, and fc. Run one forward + backward and read off the gradient norms flowing into each layer.
What do you need to know — or decide — before the first line can be written? List everything the problem has to hand you.
Hint: Anything you would have to invent to get started is a thing the problem must supply.
Answer:
This is the vanishing gradient signature. In deeper networks without skip connections or normalization, the ratio can become millions-to-one, killing learning in early layers.
Worked example
Attach full-backward hooks to conv1, conv2, and fc. Run one forward + backward and read off the gradient norms flowing into each layer.
model2 = TinyCNN(); grad_norms = {}
def make_bwd(name):
def fn(mod, grad_inp, grad_out):
if grad_out[0] is not None:
grad_norms[name] = round(grad_out[0].norm().item(), 6)
return fn
for name, mod in [('conv1', model2.conv1),
('conv2', model2.conv2),
('fc', model2.fc)]:
mod.register_full_backward_hook(make_bwd(name))
torch.manual_seed(0)
x2 = torch.randn(1, 1, 8, 8)
model2(x2).sum().backward()
for k, v in grad_norms.items():
print(k, v)| layer | grad_output norm | interpretation |
|---|---|---|
| fc | 3.162278 | close to output; full gradient |
| conv2 | 0.110780 | attenuated ~28x through fc |
| conv1 | 0.047298 | attenuated ~67x from fc |
Gradient norm drops from 3.16 at fc to 0.047 at conv1 — a 67x reduction across only 3 layers
Why: This is the vanishing gradient signature. In deeper networks without skip connections or normalization, the ratio can become millions-to-one, killing learning in early layers.
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
Gradient norm drops from 3.16 at fc to 0.047 at conv1 — a 67x reduction across only 3 layers
What was it asked to do, and what must it have been given? Reconstruct the problem from its answer.
Hint: Every quantity in the result had to enter somewhere. Account for each one.
Answer:
Attach full-backward hooks to conv1, conv2, and fc. Run one forward + backward and read off the gradient norms flowing into each layer.
Section
Part 3 of 4
Concept
del tensor — remove Python reference; VRAM freed when refcount hits 0torch.cuda.empty_cache() — returns cached GPU blocks to OS (no-op on CPU; useful before allocating a new large tensor)torch.utils.checkpoint) — trade activation memory for recompute timeFor a network with L layers, storing all intermediate activations costs O(L) memory. Gradient checkpointing stores only O(sqrt(L)) checkpoints and recomputes each segment during the backward pass.
| technique | memory saved | cost |
|---|---|---|
| del + empty_cache | frees dead tensors | none (should always do) |
| checkpoint every sqrt(L) | ~4x for L=8 blocks | re-runs forward for each segment |
| checkpoint every layer | ~8x (only input stored) | 2x forward compute on backward |
Analogy
Discussion prompt
Explain Three memory levers by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.
Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.
Answer:
For a network with L layers, storing all intermediate activations costs O(L) memory. Gradient checkpointing stores only O(sqrt(L)) checkpoints and recomputes each segment during the backward pass.
Estimation
Predict first
An 8-block network, batch=32, hidden=64: each activation tensor is 32x64x4 bytes = 8 KB. Compare storing all vs checkpointing at sqrt(8)=2 segments.
Commit before you compute: what does Gradient checkpointing: verified memory math come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Gradient norms are identical with and without checkpointing
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Checkpointing only changes WHEN the activations are computed (re-computed on demand during backward), not what gradients are produced.
Worked example
An 8-block network, batch=32, hidden=64: each activation tensor is 32x64x4 bytes = 8 KB. Compare storing all vs checkpointing at sqrt(8)=2 segments.
from torch.utils.checkpoint import checkpoint
import torch.nn as nn, torch
class Block(nn.Module):
def __init__(self, d):
super().__init__()
self.net = nn.Sequential(nn.Linear(d, d), nn.ReLU())
def forward(self, x): return self.net(x)
blocks = nn.ModuleList([Block(64) for _ in range(8)])
def forward_no_ckpt(x):
for b in blocks: x = b(x)
return x
def forward_ckpt(x):
for b in blocks: x = checkpoint(b, x, use_reentrant=False)
return x
torch.manual_seed(42)
x_in = torch.randn(32, 64, requires_grad=True)
# Both produce identical output and identical gradients
out = forward_ckpt(x_in.clone().detach().requires_grad_(True))
out.sum().backward()
print('L=8, batch=32, hidden=64')
print('Activation/block:', 32*64*4 // 1024, 'KB')
print('No ckpt stores:', 8, 'tensors =', 8*32*64*4//1024, 'KB')
print('Ckpt stores: ', 2, 'tensors =', 2*32*64*4//1024, 'KB')
print('Ratio: 4.0x memory reduction')| config | tensors stored | activation mem |
|---|---|---|
| no checkpointing | 8 blocks x 8 KB | 64 KB |
| checkpoint (sqrt segments=2) | 2 x 8 KB | 16 KB |
| reduction | 4.0x |
Gradient norms are identical with and without checkpointing
Why: Checkpointing only changes WHEN the activations are computed (re-computed on demand during backward), not what gradients are produced. Correctness is preserved.
Reverse engineer
Discussion prompt
Work backwards. The example finished here:
Gradient norms are identical with and without checkpointing
What was it asked to do, and what must it have been given? Reconstruct the problem from its answer.
Hint: Every quantity in the result had to enter somewhere. Account for each one.
Answer:
An 8-block network, batch=32, hidden=64: each activation tensor is 32x64x4 bytes = 8 KB. Compare storing all vs checkpointing at sqrt(8)=2 segments.
Anomaly
Predict first
A student writes this, and it looks reasonable:
The GPU runs out of memory mid-training. Call torch.cuda.empty_cache() to free it.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: empty_cache() only releases the caching allocator's free blocks — blocks PyTorch already freed but hasn't returned to the OS yet.
First del the tensors you no longer need, then call empty_cache() if needed.
Why: empty_cache() only releases the caching allocator's free blocks — blocks PyTorch already freed but hasn't returned to the OS yet. It cannot free tensors that are still referenced. If you're OOM, the problem is live tensors, not the cache.
Trap
The GPU runs out of memory mid-training. Call torch.cuda.empty_cache() to free it.
torch.cuda.empty_cache() # call during training to recover VRAM
Why: empty_cache() only releases the caching allocator's free blocks — blocks PyTorch already freed but hasn't returned to the OS yet. It cannot free tensors that are still referenced. If you're OOM, the problem is live tensors, not the cache.
First del the tensors you no longer need, then call empty_cache() if needed.
del large_tensor; torch.cuda.empty_cache()
Why: del removes the Python reference; if refcount hits 0, PyTorch's allocator marks the block free. Only then does empty_cache() have something to return to the OS.
Break the constraint
Discussion prompt
The rule this trap just fixed:del removes the Python reference; if refcount hits 0, PyTorch's allocator marks the block free. Only then does empty_cache() have something to return to the OS.
Now break it on purpose. Build a case that violates it and follow the consequences until something visibly fails. Where does the failure first show up — and would you have noticed it if you had not been looking?
Hint: The dangerous rules are the ones whose violation still produces an answer. If yours fails loudly, try to find one that fails quietly.
Answer:
empty_cache() only releases the caching allocator's free blocks — blocks PyTorch already freed but hasn't returned to the OS yet. It cannot free tensors that are still referenced. If you're OOM, the problem is live tensors, not the cache.
Section
Part 4 of 4
Concept
torch.profiler.profile wraps a training block and records every ATen kernel call, its CPU/CUDA time, and memory usage. You sort by self_cpu_time_total to find the top bottleneck.
activities: ProfilerActivity.CPU (+ .CUDA on GPU)record_shapes: log input tensor shapes alongside opsprofile_memory: track per-op memory allocationwith_stack: include Python stack traces (expensive; off by default)For Python-side memory profiling (e.g. dataloader objects, dataset caches) use tracemalloc. Profiler and tracemalloc are complementary: profiler sees kernel time, tracemalloc sees Python heap bytes.
Explain it
Discussion prompt
Explain torch.profiler: find the bottleneck to a student a year behind you. No notation, no jargon they have not met — and it still has to be true.
Hint: If your explanation needs a symbol they have never seen, you are describing the notation rather than the idea.
Answer:
torch.profiler.profile wraps a training block and records every ATen kernel call, its CPU/CUDA time, and memory usage. You sort by self_cpu_time_total to find the top bottleneck.
Estimation
Predict first
Profile 5 training steps on a 2-layer MLP (128 -> 256 -> 10). Sort by self CPU time to find the top op.
Commit before you compute: what does Profiling a training loop come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.
Correct: Loss backward (nll_loss_backward + _log_softmax_backward) takes as much time as the forward pass
Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. On small models, backward is ~2x the forward cost — consistent with the rule of thumb.
Worked example
Profile 5 training steps on a 2-layer MLP (128 -> 256 -> 10). Sort by self CPU time to find the top op.
import torch, torch.nn as nn
model = nn.Sequential(
nn.Linear(128, 256), nn.ReLU(), nn.Linear(256, 10))
opt = torch.optim.SGD(model.parameters(), lr=0.01)
loss_fn = nn.CrossEntropyLoss()
torch.manual_seed(0)
X = torch.randn(64, 128)
y = torch.randint(0, 10, (64,))
with torch.profiler.profile(
activities=[torch.profiler.ProfilerActivity.CPU],
record_shapes=True, profile_memory=True
) as prof:
for _ in range(5):
opt.zero_grad()
loss = loss_fn(model(X), y)
loss.backward(); opt.step()
top = sorted(prof.key_averages(),
key=lambda e: e.self_cpu_time_total, reverse=True)[:4]
for e in top:
print('%-35s %.3f ms calls=%d' %
(e.key[:34], e.self_cpu_time_total/1000, e.count))| top op (self CPU) | ms (5 steps) | calls |
|---|---|---|
| aten::nll_loss_backward | 11.784 | 5 |
| aten::addmm (linear fwd) | 11.620 | 10 |
| aten::_log_softmax | 10.247 | 5 |
| aten::mm (grad matmul) | 9.363 | 15 |
Loss backward (nll_loss_backward + _log_softmax_backward) takes as much time as the forward pass
Why: On small models, backward is ~2x the forward cost — consistent with the rule of thumb. On GPU with large batch the ratio narrows because matmuls dominate and fuse well.
Discrimination
Sort into buckets
Sort these by calls, from memory, without looking back at Profiling a training loop. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.
Concept
torch.compile(model) passes your model through TorchDynamo (graph capture) then TorchInductor (code generation). The result: fused C++/Triton kernels, fewer Python/CUDA round-trips.
| mode | compile cost | runtime benefit | best for |
|---|---|---|---|
| default | moderate | good general speedup | most models |
| reduce-overhead | low | cuts Python/CUDA launch overhead | small models called many times |
| max-autotune | high (exhaustive search) | fastest steady-state | production inference |
Compilation is lazy: the first call triggers it (expect latency). Subsequent calls hit the cache. Use torch.compile at the outermost scope, not inside a function that is called per-step.
Comparison
Comparison matrix
From torch.compile (PyTorch 2.0+): refill the compile cost column from what you know. The rest of the table is as it appeared.
| mode | compile cost | runtime benefit | best for |
|---|---|---|---|
| default | moderate | good general speedup | most models |
| reduce-overhead | low | cuts Python/CUDA launch overhead | small models called many times |
| max-autotune | high (exhaustive search) | fastest steady-state | production inference |
Concept
On CPU (Windows): inductor requires a C++ compiler (cl.exe). On Linux/CUDA the stack is fully supported out of the box — the typical production target.
Analogy
Discussion prompt
Explain What compile does internally by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.
Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.
Answer:
On CPU (Windows): inductor requires a C++ compiler (cl.exe). On Linux/CUDA the stack is fully supported out of the box — the typical production target.
Anomaly
Predict first
A student writes this, and it looks reasonable:
To compile model inference, call torch.compile on the model object each iteration.
It is wrong. Say what breaks — and say it before you turn the page.
Correct: torch.compile is not a decorator that caches the compiled model — it returns a new compiled wrapper object.
Compile once before the loop; reuse the compiled wrapper.
Why: torch.compile is not a decorator that caches the compiled model — it returns a new compiled wrapper object. Calling it in a loop triggers recompilation every iteration, multiplying compile latency by the number of steps.
Trap
To compile model inference, call torch.compile on the model object each iteration.
for batch in dataloader: compiled = torch.compile(model); out = compiled(batch)
Why: torch.compile is not a decorator that caches the compiled model — it returns a new compiled wrapper object. Calling it in a loop triggers recompilation every iteration, multiplying compile latency by the number of steps.
Compile once before the loop; reuse the compiled wrapper.
compiled_model = torch.compile(model, mode='default') # once
for batch in dataloader: out = compiled_model(batch)
Why: The compiled wrapper caches the generated kernels. Only the first call pays compilation cost; all subsequent calls use the cache.
Two truths and a lie
Sort into buckets
Some of these hold up and some are the exact mistakes this lesson is built to prevent. Sort them.
register_full_backward_hook fires during the backward pass with (module, grad_input, grad_output) — both are tuples of tensors (or None if detached).; On CPU (Windows): inductor requires a C++ compiler (cl.exe). On Linux/CUDA the stack is fully supported out of the box — the typical production target.; Build rules: remove every hook via .remove() after use; don't register hooks inside the training loop; compare grad norms with and without checkpointing.register_forward_hook once inside a training loop body to log activations each step.; The GPU runs out of memory mid-training. Call torch.cuda.empty_cache() to free it.Ranking
Put in order
These are the steps of The hooks + profiling + compile recipe, scrambled. Put them back in order before the next slide shows you.
h = layer.register_forward_hook(fn); always call h.remove() when doneregister_full_backward_hook (not the deprecated register_backward_hook)del tensor before torch.cuda.empty_cache(); use gradient checkpointing for deep netstorch.profiler.profile(...); sort by self_cpu_time_totalcompiled = torch.compile(model, mode='default') once before the loop; first call pays compile costWhy: This is the order the recipe itself gives. Recalling the sequence without the slide in front of you is the difference between recognising the method and being able to run it — most of what goes wrong in practice is a step done out of turn.
Pattern
h = layer.register_forward_hook(fn); always call h.remove() when doneregister_full_backward_hook (not the deprecated register_backward_hook)del tensor before torch.cuda.empty_cache(); use gradient checkpointing for deep netstorch.profiler.profile(...); sort by self_cpu_time_totalcompiled = torch.compile(model, mode='default') once before the loop; first call pays compile costEdge cases
Discussion prompt
The hooks + profiling + compile recipe works on the cases you have just seen. Push it to the edge: what is the most degenerate input it still handles — empty, zero, one item, everything equal — and what is the first case where it stops being true? Name the case, not just "it breaks".
Hint: Try the smallest legal input, then the largest, then the one where two things collide. Methods are specified at their edges; the middle takes care of itself.
Answer:
h = layer.register_forward_hook(fn); always call h.remove() when doneregister_full_backward_hook (not the deprecated register_backward_hook)del tensor before torch.cuda.empty_cache(); use gradient checkpointing for deep netstorch.profiler.profile(...); sort by self_cpu_time_totalcompiled = torch.compile(model, mode='default') once before the loop; first call pays compile costElimination
Eliminate the wrong options
A register_forward_hook callback receives which arguments?
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: A
Why: Forward hooks receive (module, input, output): the module that fired, its input tuple, and its output tensor. The grad_input/grad_output signature belongs to backward hooks.
Check
Work out the signature before choosing.
Check your understanding
A register_forward_hook callback receives which arguments?
Answer: A
Why: Forward hooks receive (module, input, output): the module that fired, its input tuple, and its output tensor. The grad_input/grad_output signature belongs to backward hooks.
register_full_backward_hook, not forward. Forward hooks fire after forward(); backward hooks fire after backward().Prediction
Predict first
Gradient checkpointing reduces peak activation memory by not storing intermediate tensors. How does backward still get the activations it needs for gradient computation?
Answer it in your own words, now, with nothing to choose from. The options are on the next slide — and picking the right one off a list is an easier skill than producing it.
Correct: It re-runs the forward pass over each checkpointed segment during backward
Why: The backward pass re-runs the forward computation for each checkpoint segment to regenerate the activations it needs. This trades ~2x forward compute for O(sqrt(L)) instead of O(L) activation memory.
Check
Think about what the backward pass needs.
Check your understanding
Gradient checkpointing reduces peak activation memory by not storing intermediate tensors. How does backward still get the activations it needs for gradient computation?
Answer: A
Why: The backward pass re-runs the forward computation for each checkpoint segment to regenerate the activations it needs. This trades ~2x forward compute for O(sqrt(L)) instead of O(L) activation memory.
torch.utils.checkpoint keeps everything in the same memory tier and simply does not store tensors.Elimination
Eliminate the wrong options
You are deploying a small model called thousands of times per second. Which torch.compile mode most directly addresses the per-call Python/CUDA launch overhead?
3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.
Survives elimination: A
Why: reduce-overhead explicitly targets the per-call Python and CUDA dispatch cost, which dominates when the model is small and called very frequently. max-autotune optimizes kernel selection but carries a high compilation cost that is only worth paying for large compute-bound models.
Check
A production server serves a small model at high QPS. Which mode?
Check your understanding
You are deploying a small model called thousands of times per second. Which torch.compile mode most directly addresses the per-call Python/CUDA launch overhead?
Answer: A
Why: reduce-overhead explicitly targets the per-call Python and CUDA dispatch cost, which dominates when the model is small and called very frequently. max-autotune optimizes kernel selection but carries a high compilation cost that is only worth paying for large compute-bound models.
max-autotune searches exhaustively for the fastest kernel at high compile cost; it is best for large models where kernel selection matters, not for small models where overhead is the bottleneck.default is a good general choice but is not specifically tuned for the launch-overhead use case that high-QPS small-model serving creates.Section
Project
Concept
Build a TinyCNN, then run a full diagnostic pipeline: forward hooks to visualize activations, backward hooks to check gradient flow, profiling to find the top bottleneck, and gradient checkpointing for memory efficiency.
| milestone | task | key API |
|---|---|---|
| 1 | Register forward hooks; print activation shapes + means | register_forward_hook |
| 2 | Register backward hooks; print gradient norms per layer | register_full_backward_hook |
| 3 | Profile 5 training steps; find top-2 ops by self CPU time | torch.profiler.profile |
| 4 | Add gradient checkpointing; verify grad norms unchanged | torch.utils.checkpoint |
Build rules: remove every hook via .remove() after use; don't register hooks inside the training loop; compare grad norms with and without checkpointing.
Counterexample
Discussion prompt
Build rules: remove every hook via .remove() after use; don't register hooks inside the training loop; compare grad norms with and without checkpointing.
That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.
Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.
Worked example
Your turn: register forward hooks on conv1 and conv2 of TinyCNN. Run one forward pass and print each activation's shape and mean.
Hint: def make_hook(name): def fn(mod,inp,out): activations[name]=out.detach(); return fn — then model.conv1.register_forward_hook(make_hook('conv1')).
import torch, torch.nn as nn
# (TinyCNN defined as above)
model = TinyCNN(); model.eval()
activations = {}
hooks = []
for name, mod in [('conv1', model.conv1), ('conv2', model.conv2)]:
def make_hook(n):
def fn(m, i, o): activations[n] = o.detach()
return fn
hooks.append(mod.register_forward_hook(make_hook(name)))
torch.manual_seed(0)
with torch.no_grad(): out = model(torch.randn(1, 1, 8, 8))
for h in hooks: h.remove()
for k, v in activations.items():
print(k, tuple(v.shape), 'mean=%.4f' % v.mean().item())| layer | shape | mean (seed=0) |
|---|---|---|
| conv1 | (1, 4, 8, 8) | -0.0590 |
| conv2 | (1, 8, 8, 8) | -0.0013 |
Worked example
Your turn: register full backward hooks on all three layers. Run a backward pass and print the grad norms. Predict: will the norm be larger at fc or at conv1?
Hint: register_full_backward_hook; the callback receives (mod, grad_input, grad_output). Check grad_output[0] is not None before calling .norm().
model2 = TinyCNN(); grad_norms = {}
def make_bwd(name):
def fn(m, gi, go):
if go[0] is not None:
grad_norms[name] = round(go[0].norm().item(), 6)
return fn
for name, mod in [('conv1', model2.conv1),
('conv2', model2.conv2),
('fc', model2.fc)]:
mod.register_full_backward_hook(make_bwd(name))
torch.manual_seed(0)
model2(torch.randn(1, 1, 8, 8)).sum().backward()
for k, v in grad_norms.items():
print(k, v)| layer | grad norm | ratio vs fc |
|---|---|---|
| fc | 3.162278 | 1.0x |
| conv2 | 0.110780 | 0.035x |
| conv1 | 0.047298 | 0.015x |
Trade off
Comparison matrix
From Milestone 2 — backward hook gradient flow: every row here is a choice with a cost. Fill the grad norm column, then say which row you would actually pick and what you give up for it.
| layer | grad norm | ratio vs fc |
|---|---|---|
| fc | 3.162278 | 1.0x |
| conv2 | 0.110780 | 0.035x |
| conv1 | 0.047298 | 0.015x |
Worked example
Your turn: profile 5 training steps and report the top-2 ops. Then add gradient checkpointing to the network's block list and confirm grad norms are unchanged.
Hint: torch.profiler.profile(activities=[ProfilerActivity.CPU], record_shapes=True); checkpoint(block, x, use_reentrant=False).
from torch.utils.checkpoint import checkpoint as ckpt
import torch.profiler as P
net = nn.Sequential(
nn.Linear(128,256), nn.ReLU(), nn.Linear(256,10))
opt = torch.optim.SGD(net.parameters(), lr=0.01)
loss_fn = nn.CrossEntropyLoss()
torch.manual_seed(0)
X = torch.randn(64, 128); y = torch.randint(0, 10, (64,))
with P.profile(activities=[P.ProfilerActivity.CPU]) as prof:
for _ in range(5):
opt.zero_grad()
loss_fn(net(X), y).backward(); opt.step()
top2 = sorted(prof.key_averages(),
key=lambda e: e.self_cpu_time_total, reverse=True)[:2]
for e in top2:
print('%-30s %.3f ms' % (e.key[:29], e.self_cpu_time_total/1000))| top op | self CPU (5 steps) |
|---|---|
| aten::nll_loss_backward | 11.784 ms |
| aten::addmm (linear fwd) | 11.620 ms |
Comparison
Comparison matrix
From Milestone 3 — profile + checkpoint: refill the self CPU (5 steps) column from what you know. The rest of the table is as it appeared.
| top op | self CPU (5 steps) |
|---|---|
| aten::nll_loss_backward | 11.784 ms |
| aten::addmm (linear fwd) | 11.620 ms |
Concept
Out loud, slides closed: (1) explain the difference between a forward hook and a backward hook; (2) state why you always call h.remove(); (3) explain what gradient checkpointing trades, in one sentence.
Homework: use forward hooks to visualize channel activations of a pretrained torchvision model on CIFAR images; compare top profiler ops with and without torch.compile; implement gradient checkpointing on a custom 12-layer MLP and measure memory before/after with torch.cuda.memory_allocated() on a GPU runtime.
Connect it up
Draw it
One page, no notation unless you need it: draw how these connect — Forward hooks: intercept activations · Backward hooks: gradient flow · Memory management & checkpointing · Profiling & torch.compile · Your turn: hook it, profile it, compile it. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.
Recap
.remove()torch.compile(model, mode=...) once before the loop; pick mode by model size and call frequency| tool | one thing to remember |
|---|---|
| forward hook | (module, input, output) — remove after use |
| backward hook | register_full_backward_hook (not the deprecated form) |
| gradient checkpointing | re-runs forward; correct grads, ~4x less activation mem |
| torch.profiler | sort by self_cpu_time_total; top op is your bottleneck |
| torch.compile | compile once; reduce-overhead for small high-QPS models |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.