JAX and Flax

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.

Read these first

On this page 5
  1. Why the loop looks different from other frameworks
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

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

Developer — Code and libraries.

Setup

bash
pip install flax optax

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

train_xor.py
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))
Output
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

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

What to learn next

These follow on from what you just read.

  • Natural Language Processing

    What is NLP?

    Natural language processing is how a computer reads, understands and writes human language instead of only handling numbers.

  • Natural Language Processing

    Tokenization

    Tokenization is cutting text into small pieces called tokens, because a model can only work with a fixed list of known pieces.

  • Natural Language Processing

    Embeddings

    An embedding is a list of numbers that stands for a word or a sentence, arranged so that things with similar meaning end up close together.