One hidden layer. Data: input
Chain rule, one step at a time, each reusing the last:

| reverse mode | forward mode | |
|---|---|---|
| direction | output back to inputs | inputs forward to outputs |
| one pass gives | gradient of one output w.r.t. all inputs | derivative of all outputs along one input |
| wins when | many parameters, one loss | few inputs, many outputs |
| JAX | jax.vjp, jax.grad, jacrev |
jax.jvp, jacfwd |
| PyTorch | loss.backward(), torch.func.vjp |
torch.func.jvp |
Training is reverse mode. A 3-variable simulator with 1,000 outputs is forward mode.

| PyTorch | JAX | |
|---|---|---|
| idea | record what the code did to tensors | transform the function you wrote |
| gradient | loss.backward() fills .grad |
jax.grad(f) returns a new function |
| state | inside the module and optimizer | values you pass in and get back |
| arrays | mutable | immutable |
Same numbers either way: the two agree to 4.4 × 10⁻¹⁶, float64 epsilon.

Read off a real loss.grad_fn. Rebuilt from scratch on every call: define-by-run.
.grad accumulatesloss = loss_fn(model(x), y)
loss.backward() # AccumulateGrad: W.grad += dL/dW
.grad; they do not overwrite itbackward() each, step onceoptimizer.zero_grad()Gradient accumulation: summing gradients from several backward passes before one optimizer step.
You call loss.backward() twice on the same loss (with retain_graph=True), with no zero_grad() in between. What is W.grad now?
The tape keeps every intermediate until backward() uses it. No backward(), no release.
total = 0
for x, y in val_loader:
loss = loss_fn(model(x), y)
total += loss # keeps this batch's whole graph alive
| 1,024-unit MLP, 50 validation batches | memory held |
|---|---|
total += loss |
478 MB, growing ~10 MB per batch |
total += loss.item() |
8 MB (the last batch's graph) |
inside torch.no_grad() |
0 MB |
| scenario | why | how |
|---|---|---|
| validation, test, serving | no gradient coming; the graph is pure memory | with torch.no_grad(): |
| a hand-written update | W -= lr * W.grad raises: in-place on a leaf |
inside no_grad(), as optimizer.step() does |
| part of the model frozen | pretrained layers, a target that must not move | p.requires_grad_(False), .detach() |
torch.inference_mode() is a stricter, slightly faster no_grad() for serving.
model.eval() # dropout off, batch norm uses running statistics
with torch.no_grad(): # record nothing
total = sum(loss_fn(model(x), y).item() for x, y in val_loader)
model.train() # back to training behavior
eval() and no_grad() are different switches: layers versus the tape. Evaluation needs botheval(): a validation score another script cannot reproducejax.grad, and arrays are immutable, so only freezing needs code: jax.lax.stop_gradient(x)def loss(params, x, y):
a = jnp.tanh(params["W"] @ x + params["b"])
return (params["v"] @ a + params["c"] - y) ** 2
grads = jax.grad(loss)(params, x, y) # same structure as params
.grad, nothing to zerojax.grad(jax.grad(f)) differentiates twice; jax.jit compiles it through XLAg = dot_general a e # W @ x
h = add g b # + b
i = tanh h
j = dot_general d i # v · a
k = add j c
l = sub k f # - y
m = integer_pow[y=2] l
jax.make_jaxpr(loss) prints this: 7 equations, the same 7 nodes as PyTorch's tapejax.grad rewrites it into a 16-equation program: forward and backward togetherDocs: Understanding jaxprs
if on an array's value cannot be traced: use jax.lax.condjax.lax.while_loopprint or list append runs once, at trace timex[0] = 1.0 is a TypeError, use x = x.at[0].set(1.0)PyTorch pays none of this. Define-by-run means any Python works.
vmap and jitvmap: write it for one example, map the batch axis. Agrees with a Python loop over 256 examples to 2.2 × 10⁻¹⁵.
jit on |
eager | jit | |
|---|---|---|---|
one 512×512 matmul + tanh |
2.8 ms | 3.6 ms | 0.8× |
| elementwise chain, 2M floats | 16.8 ms | 14.6 ms | 1.2× |
10-step lax.fori_loop |
83.3 ms | 30.3 ms | 2.8× |
Compilation pays where there are many small operations to fuse.
| PyTorch | JAX | |
|---|---|---|
| default float | float32 | float32 |
| float64 | on request (dtype=, set_default_dtype); not on Apple MPS |
only with jax_enable_x64; else truncated |
| batch | leading dimension in every module | vmap over a per-example function |
| compile | torch.compile, optional |
jax.jit, the normal path |
| transforms | torch.func.grad, vmap, jvp |
grad, vmap, jvp, native |
| randomness | global, torch.manual_seed |
explicit keys, jax.random.split |
| optimizers | torch.optim |
optax |
The gap is now mostly defaults: PyTorch starts eager and opts in; JAX starts functional.
Tensor: an n-dimensional array with a shape, a dtype and a device, which in PyTorch can also record the operations applied to it.
(N, features), (N, channels, time)The model returns predictions shaped (64, 1). The targets are (64,). What shape is pred - target?
(64, 1)(64,)(64, 64)model = nn.Sequential(nn.Linear(8, 64), nn.ReLU(), nn.Linear(64, 1))
model[2](x) # x: 32 mixes x 8 features, but the first layer was skipped
mat1 and mat2 shapes cannot be multiplied (32x8 and 64x1)
| piece | what it is |
|---|---|
mat1, 32x8 |
your input: 32 rows, 8 features |
mat2, 64x1 |
the layer's weight, transposed: it expects 64 features |
| the rule | inner dimensions must agree, and 8 ≠ 64 |
Fix the model, not the input: a skipped layer, or a stale in_features.
| call | dtype |
|---|---|
np.array(3.14) |
float64 |
torch.tensor(3.14) |
float32 |
torch.tensor(np.float64(3.14)) |
float64 |
jnp.asarray(np.ones(3)), x64 off (the default) |
float32, no warning |
jax.config.update("jax_enable_x64", True) before any array existsjnp.arange(3.0)[10] returns 2.0.at[10].set(v) is silently droppedRead before writing any JAX: the Sharp Bits
A story: code that ran, a number that looked fine, then the cause. Three today.
How it looked
tensor.dtype printed: float64PyTorch is off by 7.5e-8, JAX by 4.4e-16, and every dtype prints float64. Most likely cause?
1.234, and a bare rng.normal()torch.tensor stored both as float32dtype=torch.float64 on both: 1.7 × 10⁻¹⁸Check dtypes where data enters, not after. An error near 10⁻⁷ means float32.
Predict concrete strength from the mix and its age: 1,030 rows, 8 inputs, a crush test that takes weeks
| any real model must beat | a working MLP gets |
|---|---|
| 17.9 MPa: predict the training mean for every row | 5.5 MPa |
folds = KFold(5, shuffle=True, random_state=0).split(X)
for seed in range(5):
for tr, va in folds:
tree = HistGradientBoostingRegressor(random_state=seed).fit(X[tr], y[tr])
net = train(X[tr], y[tr], X[va], y[va], seed=seed)
KFold, 75% of validation rows share a mix with training: the same concrete at another ageGroupKFold by mix: tree 6.25, MLP 6.48, a tie; the leak was worth 1.81 MPa to the tree, 1.61 to the net
The textbook leak: StandardScaler fitted on all rows before splitting (Lecture 7)
| scaler fitted on | RMSE, 3 seeds × 5 folds |
|---|---|
| training rows only | 6.27 MPa |
| all rows (leaky) | 6.23 MPa |
| difference | −0.04 ± 0.06 |
for xb, yb in loader:
loss = loss_fn(model(xb), yb)
loss.backward()
opt.step()

No opt.zero_grad(): .grad sums every past step, so each update follows the sum of all past gradients. Ends at 23.2 MPa, worse than the mean. Fixed: 5.5.
| story | looked like | caught by |
|---|---|---|
| dtype | a library difference | an exact reference, 10⁻⁷ fingerprint |
| leaky split | trees beat nets | a split grouped by mix |
no zero_grad |
a noisy learning rate | reading the loop |
target shape (N,) |
a weak first model | the mean baseline, prediction spread |
| raw inputs + Adam | a respectable model | the input scales, or trying SGD |
None raised an error. All four models, rerun and taken apart: l11-four-models.ipynb
for epoch in range(n_epochs):
model.train()
for xb, yb in loader:
xb, yb = xb.to(device), yb.to(device)
loss = loss_fn(model(xb), yb) # forward
optimizer.zero_grad() # clear the accumulator
loss.backward() # backward
optimizer.step() # update
DataLoader, nn.Module, loss function, optimizer. Story 3 was one missing line here@jax.jit
def step(params, opt_state, xb, yb):
loss, grads = jax.value_and_grad(loss_fn)(params, xb, yb)
updates, opt_state = optimizer.update(grads, opt_state, params)
return optax.apply_updates(params, updates), opt_state, loss
zero_grad, no .to(device): state goes in and comes outAdam: momentum on the gradient, divided by a running root-mean-square of the gradient, so each parameter gets its own step size.

Narrow bowl, 40× steeper in

Gradient drops 100× at step 200: with
torch.optim |
optax |
|
|---|---|---|
| Adam learning rate | 1e-3 | required |
| 0.9, 0.999, 1e-8 | 0.9, 0.999, 1e-8 | |
| AdamW weight decay | 0.01 | 0.0001 |
SGD on the same fold, 120 epochs:
| lr | 0.001 | 0.01 | 0.1 | 1.0 | 2.0 |
|---|---|---|---|---|---|
| RMSE, MPa | 11.9 | 6.4 | 5.1 | 91 | nan from epoch 1 |
nan in two, 91 to 1,235 MPa in the other fourmodel.to(device), x.to(device); mismatch is a loud error, the good casejax.devices() shows it.cpu() or print inside the loop forces a sync and serializes everythingtorch.cuda.synchronize() or x.block_until_ready()Today's MLP (two hidden layers of 64, 1,030 rows) moves from the laptop CPU to its GPU. Time per epoch?

Today's model: 8.6 ms per epoch on the CPU, 21.1 ms on the GPU. The GPU wins past a few hundred hidden units.
| PyTorch | JAX | |
|---|---|---|
| strongest at | standard architectures, pretrained models, deployment | differentiating simulators, ODE solves, physical models |
| composes | eager code, plus torch.func and torch.compile |
vmap(grad(f)), jax.hessian, one line each |
| costs you | hidden state: .grad, train()/eval(), global seed |
purity, lax.cond, tracing and recompiles |
| fails quietly with | missing zero_grad, float32 from a Python float |
float32 by default, clamped indices |
And on 1,030 rows of concrete, gradient boosting ties the MLP in under a second, with no learning rate.
l11-tensors-autograd.ipynbA gradient by hand, checked against backward() and jax.grad. A loop by hand, broken three ways. The same loop on the GPU. Net against tree, under both splits.
.grad; JAX transforms pure functionsNicknames only. Everyone who skipped one still counted in every bar you saw.
Practice module for this session, for participation credit
Reading PyTorch: Learn the Basics, JAX Sharp Bits, Grinsztajn et al. 2022
Full notes, with all sources: lectures/l11/notes.md
90 minutes of deck, then 20 of notebook and questions. Budget: AD 12, PyTorch/JAX 18, tensors 14, stories 18, loop + Adam 15, devices 8, recap 5. Four clicker questions: slides marked "a question". Each takes about 3 minutes with the re-vote. Dataset all session: UCI concrete compressive strength, 1,030 rows, 8 inputs, MPa out. If running long, cut the Adam-steps slide and the side-by-side table, never a story.
Open with the scale problem. The answer is the chain rule, done by bookkeeping.
Write these on the board if there is one. Point out that 1 - a^2 needs a, the stored forward value.
Walk the red arrows right to left from dL/dL = 1: gradient arriving from the right, times the factor on the edge. Two things to say: 1. One backward walk gives every parameter's gradient, because the loss is a scalar. 2. The backward pass needs a and x from the forward pass. That is why training uses more memory than inference.
In deep learning, reverse mode is called backpropagation. Baydin et al. 2018 has the history: https://arxiv.org/abs/1502.05767
The V: truncation error on the right, round-off on the left. You only know where the bottom is because the exact answer was available. PyTorch 1.7e-18, JAX 4.4e-16 (x64). Leave the red "careless" line for now; it is the first story.
The rest of this section is the consequence of the first row.
Left: seven nodes in the order the forward pass made them, and what each saved. Right: the graph backward() walks. Point at AccumulateGrad at the bottom: += into .grad. x and y get no nodes because nobody asked for their gradient.
retain_graph is there so D is wrong for the right reason. Without it the second call errors because the saved tensors were freed, which is a good aside if someone asks.
Measured on this laptop's MPS GPU, torch 2.7. One forward pass at batch 8,192 holds 67 MB with the tape, 0.03 MB without. The session's 64-unit model holds 0.8 MB, which is why nobody notices on a small problem: this is the out-of-memory crash halfway through an epoch on a real one.
The in-place error text: "a leaf Variable that requires grad is being used in an in-place operation." If it were allowed, the update itself would be recorded as part of the model.
Students conflate these two constantly. no_grad is about the tape; eval is about layer behavior.
Same network as the graph slide. grads is a dict with W, b, v, c.
A single matmul is already one BLAS call; jit only adds dispatch. Wall-clock on a laptop: a rerun gave 1.2x and 4.3x for the ends. Trust the order.
The broadcasting rule is on the previous slide. Expect a lot of D. This is the bug behind a model in the notebook that learned a constant: nn.MSELoss on these shapes averages every prediction minus every target.
nn.Linear(64, 1) stores weight (1, 64) and computes x @ W.T, which is why mat2 reads 64x1. Batch 32 on purpose: with a batch of 64 the two 64s read alike and the message is much harder to parse. The tempting wrong fix is x.reshape(...) until it runs.
This happened while the notes were being written. The figure was going to say "both match to machine precision".
D is ruled out: it ran on the CPU. C is ruled out because JAX matched the same reference.
The red line on the finite-difference figure is this bug. Checking tensor.dtype at the end would never have found it.
Keep 17.9 in view: it is the number every broken model gets compared with. The mix structure is what story 2 turns on. Numbers: fold 0 of a GroupKFold by mix (Lecture 9).
Give them a minute. The answer is the first line, and nothing in the output points at it.
Nothing about either model changed, only the split. Part of "trees beat nets on small tabular data" was, here, a statement about the split. A single-seed version of this showed the MLP winning. Five seeds: a tie.
I expected the scaler to be the story. Measured, it is noise here. Say that honestly.
Four lines. The missing one is opt.zero_grad(). The clicker earlier was the same fact.
Left panel. Where it ends depends on where you stop; several epochs earlier it sat near 12. The notebook prints the size of .grad step by step with and without zeroing. Middle and right panels come back in the loop section.
The last two rows are in the notebook only: target shape 17.7 -> 5.5 MPa with predictions spread 0.26 MPa and 1,560 warnings; Adam on raw inputs 8.3 -> 5.5, SGD on the same inputs NaN in the first epoch.
Left: SGD zig-zags, momentum curls, Adam walks diagonally. This per-coordinate rescaling is why Adam survived raw inputs in the notebook's fourth model (8.3 MPa, where SGD gave NaN), and why it hid the bug. Middle: beta1 = 0.99 overshoots, loss 5.3 after 100 steps vs 0.0011. Right: lr = 0.01 is still 1.6 away after 300 steps. Rotate the bowl 45 degrees and Adam's loss goes from 0.0011 to 1.1: it rescales axes, it cannot unrotate.
Bias correction: without it the first step is about 3.2 times the learning rate. Cut this slide if short on time.
Apple MPS on a laptop. A datacenter CUDA card moves the crossover and raises the plateau; it does not remove the fixed cost. Minimum over repeated trials, because interference only ever makes a timing slower. Debug on CPU with a tiny subset, then launch the real run on the GPU.
The last 20 minutes, notebook then questions. Ask for a prediction before the net-vs-tree cell prints.
Skip this slide if no clicker questions were run.