A full training loop in JAX
Everything this section built — arrays, grad, jit, pytrees, keys, Flax, Optax — assembles into one compiled training loop under 50 lines.
- 8 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
A JAX training loop is a relay race: parameters and optimizer state are batons, passed hand to hand through one compiled step, over and over.
Think of the relay properly. Nothing about the runners changes between laps — the same step repeats. What changes is what they carry: each lap hands over slightly better parameters and an updated coach's notebook. The runner is the train step; the batons are the state.
Every previous lesson in this section built one piece of this race. This lesson runs it.
Why the loop looks different from other frameworks
In Keras you call fit and a loop happens somewhere inside. In JAX you write the loop yourself, and it stays short. Each piece is one line: the model computes, grad differentiates, the optimizer prescribes, the update applies.
The reward for writing it yourself: the entire step — forward pass, gradients, optimizer arithmetic, everything — is compiled into a single fast program. And nothing is hidden. When training misbehaves at 2 a.m., every value in the loop is a plain variable you can print.
How it works
┌────────────────── one compiled train step ──────────────────┐
│ │
params ─┤─▶ predict ─▶ measure error ─▶ blame each knob ─▶ prescribe ─┤─▶ new params
state ─┤ (optimizer) ├─▶ new state
│ │
└─────── repeat, feeding each step's outputs to the next ─────┘The loop's whole body is: hand in the batons, receive improved batons. When the laps finish, the final parameters are the trained model.
A real example you have seen
Learning to cycle. Each attempt: ride (predict), wobble or fall (measure error), sense what went wrong (blame), adjust grip and balance (update). The rider after 200 laps holds nothing extra — the improvement is the new state of the rider. Training a model is that loop, made explicit enough to print.
Remember this
- The loop = one step repeated, with params and optimizer state passed through.
- Each step: predict → error → gradients → update — every arrow from this section.
- Compile the whole step once; the loop itself stays plain Python.
What to learn next
- Backpropagation — the algorithm inside
value_and_grad, in full. - Building models with Flax — revisit the module layer now that you have seen it in context.
- How neural networks learn — the same loop, told from the theory side.
Developer — Code and libraries.
Setup
pip install flax optaxThese pull in JAX. Outputs verified with jax 0.6.2, flax 0.10.7, optax 0.2.8, CPU. Seeded, so this run reproduces; across versions the exact losses may drift.
The whole thing
XOR as the dataset: four rows, unsolvable by any straight line, so the network must genuinely learn — and small enough that every number can be watched.
import jax
import jax.numpy as jnp
import flax.linen as nn
import optax
class TinyNet(nn.Module):
@nn.compact
def __call__(self, x):
x = nn.Dense(16)(x)
x = nn.relu(x)
return nn.Dense(1)(x)
# XOR: the smallest dataset a linear model cannot solve
x = jnp.array([[0., 0.], [0., 1.], [1., 0.], [1., 1.]])
y = jnp.array([[0.], [1.], [1.], [0.]])
model = TinyNet()
params = model.init(jax.random.key(0), x)
optimizer = optax.adam(0.05)
opt_state = optimizer.init(params)
def loss_fn(params, x, y):
logits = model.apply(params, x)
return jnp.mean(optax.sigmoid_binary_cross_entropy(logits, y))
@jax.jit
def train_step(params, opt_state, x, y):
loss, grads = jax.value_and_grad(loss_fn)(params, x, y)
updates, opt_state = optimizer.update(grads, opt_state, params)
params = optax.apply_updates(params, updates)
return params, opt_state, loss
for step in range(401):
params, opt_state, loss = train_step(params, opt_state, x, y)
if step % 100 == 0:
print(f"step {step:3d} loss {loss:.4f}")
preds = jax.nn.sigmoid(model.apply(params, x))
print("predictions:", jnp.round(preds.ravel(), 2))step 0 loss 0.6802 step 100 loss 0.0013 step 200 loss 0.0005 step 300 loss 0.0003 step 400 loss 0.0002 predictions: [0. 1. 1. 0.]
XOR: learned. Every lesson of this section is in the file — find them.
The walkthrough
train_step is the whole story. State in — (params, opt_state) — state out, plus the loss for logging. It is a pure function: call it twice with the same inputs and get identical outputs. That purity is what lets @jax.jit compile the entire step — forward pass, backward pass, and Adam's arithmetic fuse into one program. The first call traces and compiles; the remaining 400 replay.
value_and_grad delivers the loss (for the log) and gradients (for the update) from one shared pass — the training-loop reason that function exists.
The model outputs logits — raw scores, no sigmoid — and optax.sigmoid_binary_cross_entropy folds the sigmoid into the loss, the numerically safe arrangement. Sigmoid gets applied explicitly only at prediction time, at the bottom.
The Python loop stays outside jit — fine, because each iteration is one compiled call, and the print is honest Python. Squeezing further (the loop itself as lax.scan, zero Python per epoch) is the standard next optimisation for big models; at this size it would buy nothing.
Reassignment is the baton pass. params, opt_state, loss = train_step(...) — nothing mutates; names rebind to new trees. Forget to catch the returns and training silently reprocesses the same initial state 400 times.
Scaling this skeleton up is remarkably direct: real projects add batching over a dataset, a key split per step for dropout, checkpointing of (params, opt_state) — but the five lines inside train_step survive almost unchanged into serious codebases (often bundled into a TrainState object from flax.training).
Common mistakes
Dropping a returned baton. Catching params but not opt_state restarts Adam's memory every step — training crawls and nobody crashes. Audit the tuple.
Jitting a step that closes over changing values. train_step reads model and optimizer from the enclosing scope — safe, because both are fixed structure. Closing over something you change between calls (a Python counter, a flag) bakes the trace-time value in forever. Changing things enter as arguments.
Logging every step. loss is a device array; printing it forces a wait for the computation (async dispatch). Printing every step can serialise the pipeline — log every N steps, as the loop above does.
Recompiling per shape. Feed batches of varying size and each new shape retraces the step. Fixed batch shapes, pad the last batch — the jit rules never stop applying.
Try it yourself
Break it on purpose, twice. First: swap optax.adam(0.05) for optax.sgd(0.05) and watch step-400 loss — momentumless descent on a non-convex problem. Second: restore Adam but change the init key to jax.random.key(7) — same architecture, different starting point, different (probably still successful) path. You now know which knob did what.
What to learn next
- Backpropagation — the algorithm inside
value_and_grad, in full. - Building models with Flax — revisit the module layer now that you have seen it in context.
- How neural networks learn — the same loop, told from the theory side.
Researcher — Mathematics and papers.
The loop as a dynamical system
Training is iteration of one map: $z_{t+1} = U(z_t; B_t)$ where $z_t = (\theta_t, s_t)$ bundles parameters and optimizer state, and $B_t$ is the step's data batch — here the full four-row dataset, making this full-batch gradient descent rather than SGD; with $B_t$ sampled, the iteration becomes the stochastic approximation process of Robbins and Monro (1951). The JAX contribution is that $U$ is reified: a pure, jaxpr-compiled function, not a code path smeared across framework internals. Properties follow — determinism given $(z_0, {B_t})$, bitwise-replayable runs, and single-object state checkpointing (serialise $z_t$, a pytree).
XOR's role is historically pointed: Minsky and Papert (1969, Perceptrons) used its linear non-separability against single-layer networks, and its solution by hidden layers plus backpropagation (Rumelhart, Hinton, Williams 1986) is the origin story of deep learning. The 16-unit layer here is far wider than the minimal 2-unit solution — overparameterisation chosen deliberately: optimisation landscapes of wider networks contain fewer bad local minima for gradient methods, an empirical regularity with partial theory (Choromanska et al. 2015; the NTK regime of Jacot et al. 2018 as the extreme case).
Whole-step compilation, quantified
Fusing forward + backward + optimizer into one XLA program eliminates per-op dispatch and enables cross-boundary optimisation: gradient computation reuses forward intermediates in-buffer (liveness analysis over the whole step), Adam's elementwise chain fuses into few kernels, and — on accelerators — the step launches as one executable. The measurable consequences at scale: step time dominated by FLOPs rather than overhead, and memory high-water marks predictable from the jaxpr. The remaining Python cost — the loop and logging — amortises as $O(1)$ per step and vanishes entirely under lax.scan-based epochs, the pattern used by production LLM training stacks built on this exact skeleton (T5X, MaxText lineage).
What the skeleton omits, and where it goes: data pipelines (host-side, feeding device batches asynchronously), multi-device parallelism (jax.pmap historically, sharding-based jit currently — the same pure step, partitioned), and checkpointing (Orbax, serialising $z_t$). None alter the five-line core.
References
- Rumelhart, Hinton, Williams (1986), Learning representations by back-propagating errors.
- Robbins and Monro (1951), A stochastic approximation method.
- Jacot, Gabriel, Hongler (2018), Neural tangent kernel: convergence and generalization in neural networks.
What to learn next
- Backpropagation — the algorithm inside
value_and_grad, in full. - Building models with Flax — revisit the module layer now that you have seen it in context.
- How neural networks learn — the same loop, told from the theory side.