Lesson 70: PyTorch Hooks, Profiling & torch.compile

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

What this lesson covers

The lesson, slide by slide

1. PyTorch Hooks, Profiling & torch.compile

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.

2. By the end of this lesson you can

Objectives

  1. Register register_forward_hook to capture any layer's output activations
  2. Register register_full_backward_hook to inspect and modify gradient flow
  3. Apply gradient checkpointing (torch.utils.checkpoint) to cut activation memory by up to 4x
  4. Use torch.profiler and tracemalloc to identify the top CPU bottleneck in a training loop
  5. Call torch.compile(model, mode=...) and select the right mode for inference vs training

3. What survived from kNN and the Curse of Dimensionality?

Warm-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.

4. Forward hooks: intercept activations

Section

Part 1 of 4

5. What a forward hook is

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 typesignaturetypical 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

6. Fill in: signature for What a forward hook is

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 typesignaturetypical 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

7. Guess the shape of the answer: Hooking a CNN to capture activations

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.

8. Hooking a CNN to capture activations

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())
layeractivation shapemean (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.

9. What each one costs: Hooking a CNN to capture activations

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.

layeractivation shapemean (seed=0)
conv1(1, 4, 8, 8)-0.0590
conv2(1, 8, 8, 8)-0.0013
output logits(1, 10)argmax=9

10. Something is wrong here: forgetting to remove hooks

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.

11. Trap: forgetting to remove hooks

Trap

The 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.

The fix

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.

12. Backward hooks: gradient flow

Section

Part 2 of 4

13. register_full_backward_hook

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.

14. Break it if you can: register_full_backward_hook

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.

15. What has to be given first: Gradient flow table via backward hooks

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.

16. Gradient flow table via backward hooks

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)
layergrad_output norminterpretation
fc3.162278close to output; full gradient
conv20.110780attenuated ~28x through fc
conv10.047298attenuated ~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.

17. Work backwards from the answer: Gradient flow table via backward hooks

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.

18. Memory management & checkpointing

Section

Part 3 of 4

19. Three memory levers

Concept

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.

techniquememory savedcost
del + empty_cachefrees dead tensorsnone (should always do)
checkpoint every sqrt(L)~4x for L=8 blocksre-runs forward for each segment
checkpoint every layer~8x (only input stored)2x forward compute on backward

20. By analogy: Three memory levers

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.

21. Guess the shape of the answer: Gradient checkpointing: verified memory math

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.

22. Gradient checkpointing: verified memory math

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')
configtensors storedactivation mem
no checkpointing8 blocks x 8 KB64 KB
checkpoint (sqrt segments=2)2 x 8 KB16 KB
reduction4.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.

23. Work backwards from the answer: Gradient checkpointing: verified memory math

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.

24. Something is wrong here: calling empty_cache() to free live tensors

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.

25. Trap: calling empty_cache() to free live tensors

Trap

The 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.

The fix

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.

26. Break it on purpose: calling empty_cache() to free live tensors

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.

27. Profiling & torch.compile

Section

Part 4 of 4

28. torch.profiler: find the bottleneck

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.

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.

29. Teach it back: torch.profiler: find the bottleneck

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.

30. Guess the shape of the answer: Profiling a training loop

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.

31. Profiling a training loop

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_backward11.7845
aten::addmm (linear fwd)11.62010
aten::_log_softmax10.2475
aten::mm (grad matmul)9.36315

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.

32. Which is which, by calls

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.

5
aten::nll_loss_backward; aten::_log_softmax
10
aten::addmm (linear fwd)
15
aten::mm (grad matmul)
g1
calls is "5" for aten::nll_loss_backward, aten::_log_softmax — that is what the table on "Profiling a training loop" records, and it is the single property separating this group from the rest.
g2
calls is "10" for aten::addmm (linear fwd) — that is what the table on "Profiling a training loop" records, and it is the single property separating this group from the rest.
g3
calls is "15" for aten::mm (grad matmul) — that is what the table on "Profiling a training loop" records, and it is the single property separating this group from the rest.

33. torch.compile (PyTorch 2.0+)

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.

modecompile costruntime benefitbest for
defaultmoderategood general speedupmost models
reduce-overheadlowcuts Python/CUDA launch overheadsmall models called many times
max-autotunehigh (exhaustive search)fastest steady-stateproduction 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.

34. Fill in: compile cost for torch.compile (PyTorch 2.0+)

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.

modecompile costruntime benefitbest for
defaultmoderategood general speedupmost models
reduce-overheadlowcuts Python/CUDA launch overheadsmall models called many times
max-autotunehigh (exhaustive search)fastest steady-stateproduction inference

35. What compile does internally

Concept

  1. TorchDynamo traces the Python bytecode and builds an FX graph — a static computational graph free of Python control flow
  2. TorchInductor lowers the FX graph to optimized kernels: operator fusion (Linear + ReLU as one pass), tiling for cache locality
  3. Compiled kernels are cached; guard conditions re-trigger compilation if input shapes change
  4. Fallback: if a graph break occurs (unsupported op, dynamic shape), that segment runs eagerly

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.

36. By analogy: What compile does internally

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.

37. Something is wrong here: compiling inside the training loop

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.

38. Trap: compiling inside the training loop

Trap

The 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.

The fix

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.

39. Which of these survive contact with Lesson 70: PyTorch Hooks, Profiling &…?

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.

Holds up
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.
Breaks
Register 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.
sound
These are stated as this lesson states them — each one survives the edge cases Lesson 70: PyTorch Hooks, Profiling & torch.compile puts it through.
flawed
Each of these is lifted from a trap in this deck: reasonable-sounding, and wrong in a way that only shows up once you rely on it.

40. Rebuild the recipe: The hooks + profiling + compile recipe

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.

  1. Forward hook: h = layer.register_forward_hook(fn); always call h.remove() when done
  2. Backward hook: register_full_backward_hook (not the deprecated register_backward_hook)
  3. Memory: del tensor before torch.cuda.empty_cache(); use gradient checkpointing for deep nets
  4. Profile: wrap training steps in torch.profiler.profile(...); sort by self_cpu_time_total
  5. Compile: compiled = torch.compile(model, mode='default') once before the loop; first call pays compile cost

Why: 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.

41. The hooks + profiling + compile recipe

Pattern

  1. Forward hook: h = layer.register_forward_hook(fn); always call h.remove() when done
  2. Backward hook: register_full_backward_hook (not the deprecated register_backward_hook)
  3. Memory: del tensor before torch.cuda.empty_cache(); use gradient checkpointing for deep nets
  4. Profile: wrap training steps in torch.profiler.profile(...); sort by self_cpu_time_total
  5. Compile: compiled = torch.compile(model, mode='default') once before the loop; first call pays compile cost

42. Where does it stop working: The hooks + profiling + compile recipe

Edge 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:

  1. Forward hook: h = layer.register_forward_hook(fn); always call h.remove() when done
  2. Backward hook: register_full_backward_hook (not the deprecated register_backward_hook)
  3. Memory: del tensor before torch.cuda.empty_cache(); use gradient checkpointing for deep nets
  4. Profile: wrap training steps in torch.profiler.profile(...); sort by self_cpu_time_total
  5. Compile: compiled = torch.compile(model, mode='default') once before the loop; first call pays compile cost

43. Rule out three: Check yourself — forward hooks

Elimination

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.

  • A. (module, input, output)
  • B. (module, grad_input, grad_output)
  • C. (loss, output, target)
  • D. (input, output)

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.

44. Check yourself — forward hooks

Check

Work out the signature before choosing.

Check your understanding

A register_forward_hook callback receives which arguments?

  • A. (module, input, output) (correct)
  • B. (module, grad_input, grad_output)
  • C. (loss, output, target)
  • D. (input, output)

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.

Why B tempts people
That is the signature of register_full_backward_hook, not forward. Forward hooks fire after forward(); backward hooks fire after backward().
Why C tempts people
No PyTorch hook receives the loss or target directly; hooks are attached to layers, not to loss functions.
Why D tempts people
The module itself is always the first argument — hooks need it to identify which layer fired, especially when the same function is registered on multiple layers.

45. Answer it before you see the options: Check yourself — gradient checkpointing

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.

46. Check yourself — gradient checkpointing

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?

  • A. It re-runs the forward pass over each checkpointed segment during backward (correct)
  • B. It stores all activations on CPU instead of GPU
  • C. It approximates the activations using the input only
  • D. It skips gradient computation for checkpointed layers

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.

Why B tempts people
CPU offloading is a separate technique (e.g., Fairscale's offload). Vanilla torch.utils.checkpoint keeps everything in the same memory tier and simply does not store tensors.
Why C tempts people
Approximations would produce wrong gradients — PyTorch never approximates. The recomputed values are exact.
Why D tempts people
Skipping gradients would break backprop entirely. All layers still receive correct gradients; only the timing of activation computation changes.

47. Rule out three: Check yourself — torch.compile mode

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.

  • A. reduce-overhead
  • B. max-autotune
  • C. default
  • D. eager (no compile)

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.

48. Check yourself — torch.compile mode

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?

  • A. reduce-overhead (correct)
  • B. max-autotune
  • C. default
  • D. eager (no compile)

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.

Why B tempts people
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.
Why C tempts people
default is a good general choice but is not specifically tuned for the launch-overhead use case that high-QPS small-model serving creates.
Why D tempts people
Without compile, every call pays full Python dispatcher overhead. For thousands of calls per second, that overhead adds up to significant latency.

49. Your turn: hook it, profile it, compile it

Section

Project

50. Project: full diagnostic pipeline

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.

milestonetaskkey API
1Register forward hooks; print activation shapes + meansregister_forward_hook
2Register backward hooks; print gradient norms per layerregister_full_backward_hook
3Profile 5 training steps; find top-2 ops by self CPU timetorch.profiler.profile
4Add gradient checkpointing; verify grad norms unchangedtorch.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.

51. Break it if you can: Project: full diagnostic pipeline

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.

52. Milestone 1 — forward hook activation capture

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())
layershapemean (seed=0)
conv1(1, 4, 8, 8)-0.0590
conv2(1, 8, 8, 8)-0.0013

53. Milestone 2 — backward hook gradient flow

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)
layergrad normratio vs fc
fc3.1622781.0x
conv20.1107800.035x
conv10.0472980.015x

54. What each one costs: Milestone 2 — backward hook gradient flow

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.

layergrad normratio vs fc
fc3.1622781.0x
conv20.1107800.035x
conv10.0472980.015x

55. Milestone 3 — profile + checkpoint

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 opself CPU (5 steps)
aten::nll_loss_backward11.784 ms
aten::addmm (linear fwd)11.620 ms

56. Fill in: self CPU (5 steps) for Milestone 3 — profile + checkpoint

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 opself CPU (5 steps)
aten::nll_loss_backward11.784 ms
aten::addmm (linear fwd)11.620 ms

57. Show it off

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.

58. Connect it up: Lesson 70: PyTorch Hooks, Profiling & torch.compile

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.

59. What you can do now

Recap

toolone thing to remember
forward hook(module, input, output) — remove after use
backward hookregister_full_backward_hook (not the deprecated form)
gradient checkpointingre-runs forward; correct grads, ~4x less activation mem
torch.profilersort by self_cpu_time_total; top op is your bottleneck
torch.compilecompile once; reduce-overhead for small high-QPS models

Sources

  1. USAAIO Year-Long Master Lesson Plan, Lesson 70 (hooks, memory, profiling, torch.compile) — Barron · USAAIO Round 2 Preparation, 2026
  2. Forward/backward hooks, gradient checkpointing, torch.profiler, torch.compile API — torch 2.7.1+cpu, real execution verified June 2026

Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.

Book on Wyzant · Text (657) 465-8108