JAX and Flax

Building models with Flax

Flax gives JAX its layers — modules that describe the architecture while the learned numbers live outside, in a pytree you hold yourself.

On this page 5
  1. Why the separation
  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.

Flax is the layers library for JAX — with one twist: the model describes the architecture, while the learned numbers live outside it, in a container you hold.

Think of a cookie mould and the dough. The mould fixes the shape; it stamps but contains nothing. The dough is what actually becomes the cookie. In PyTorch or Keras, mould and dough travel as one object — the model owns its weights. In Flax, the module is the mould only. The parameters — the learned numbers — sit in a separate container you pass in at every call.

Why the separation

It follows from everything JAX is. Transformations work on pure functions: values in, values out, nothing hidden. A model object that secretly owns changing weights is exactly the hiddenness JAX forbids.

So Flax splits the two roles. The module is fixed and describable — safe to trace and compile. The parameters are data — a pytree. So grad can differentiate with respect to it, optimizers can update it, and checkpoints can save it, all with the standard tree machinery. Nothing about training a Flax model needs new concepts; the parameters are one more tree.

How it works

define:   module = the recipe of layers          (no numbers inside)

init:     module + key + example input ──▶ params     (the dough, made once)

apply:    module + params + real input ──▶ output     (every actual call)

Two verbs to own. init runs the model once on example data to discover every weight's shape, then returns freshly created parameters. This is why it needs a random key: starting values are random. apply is the actual forward pass: same module, whichever parameters you hand it.

A real example you have seen

A government form and its filled copies. The blank form (module) is printed once and never changes. Each filled copy (params) is separate paper. The clerk processes form plus filled copy together — and can process a thousand different copies through one form. Holding parameters outside the model is what lets JAX train many copies, average them, or swap them freely.

Remember this

  • A Flax module holds architecture only — no learned numbers inside.
  • init creates the parameter tree; apply runs the model with it.
  • Parameters are a plain pytree — everything from the previous lessons applies to them.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install flax

This pulls in JAX. Outputs verified with jax 0.6.2 and flax 0.10.7, CPU. We use flax.linen — the API you will meet in most existing codebases; Flax's newer NNX API changes the style (mutable Python objects) and is the project's recommended direction for new code, so expect to meet both. Linen concepts transfer.

Define, init, apply

tiny_mlp.py
import jax
import jax.numpy as jnp
import flax.linen as nn


class TinyMLP(nn.Module):
    hidden: int          # hyperparameters are class fields, not __init__ args

    @nn.compact
    def __call__(self, x):
        x = nn.Dense(self.hidden)(x)
        x = nn.relu(x)
        x = nn.Dense(1)(x)
        return x


model = TinyMLP(hidden=8)

x = jnp.ones((4, 3))                          # a fake batch: 4 rows, 3 features
params = model.init(jax.random.key(0), x)     # params live OUTSIDE the model

print(jax.tree.map(lambda p: p.shape, params))

out = model.apply(params, x)                  # every call takes params explicitly
print("output shape:", out.shape)
print("param count:", sum(p.size for p in jax.tree.leaves(params)))
Output
{'params': {'Dense_0': {'bias': (8,), 'kernel': (3, 8)}, 'Dense_1': {'bias': (1,), 'kernel': (8, 1)}}}
output shape: (4, 1)
param count: 41

The walkthrough

hidden: int as a class field — linen modules are dataclasses. Configuration goes in fields; no __init__, no super().__init__(). The module instance is immutable, which is what makes it safe scenery for jit.

@nn.compact lets you declare sublayers inline, right where they are used — nn.Dense(self.hidden)(x) creates-and-connects in one motion. The alternative style (a setup method naming each sublayer) exists for models whose parts are called from several places.

model.init(key, x) never trains anything — it runs the forward pass in shape-discovery mode. The first Dense sees 3 features arrive, so its kernel becomes (3, 8): shapes are inferred from the example input, never declared. The key seeds the random starting values, whose distribution matters more than beginners expect — that story is weight initialisation.

The returned tree is readable. {'params': {'Dense_0': {...}}} — named by layer, exactly the pytree shape audit from the last lesson. Count check: 3×8+8 + 8×1+1 = 41.

model.apply(params, x) is the forward pass. Nothing was stored by init; hand apply a different compatible tree and it computes with that instead — the swap-the-dough freedom the beginner block promised. Gradients follow as ordinary grad:

python
def loss_fn(params, x, y):
    return jnp.mean((model.apply(params, x) - y) ** 2)

grads = jax.grad(loss_fn)(params, x, jnp.ones((4, 1)))

Common mistakes

Calling the module directly. model(x) raises — an unbound module has no parameters to compute with. Forward passes go through apply. (Inside __call__, calling sublayers directly is fine; Flax binds them there.)

Losing the 'params' wrapper. init returns {'params': {...}}, and apply expects the same wrapper. Extracting the inner dict and passing it bare produces scope errors. Keep the wrapper; index inside it only for inspection.

Expecting weights on the object. model.hidden exists; model.weights does not. If you find yourself hunting for parameters inside the module, the mental model has slipped — they are in your params variable.

Fresh init keys every run. init(jax.random.key(time()), ...) makes runs unrepeatable. Fix the seed, thread keys deliberately — the PRNG discipline applies from the very first parameter.

Try it yourself

Add a second hidden layer and a dropout_rate: float field with nn.Dropout(self.dropout_rate, deterministic=True)(x) between layers. Re-init, reprint the shape tree, and recount parameters by hand before checking against jax.tree.leaves. Then flip deterministic to False and see what apply starts demanding — an rngs={'dropout': key} argument: named randomness, threaded explicitly.

What to learn next

Researcher — Mathematics and papers.

Modules as parameter-tree factories

A linen module is a lens over a pure function: apply(variables, x) evaluates $f(\theta, x)$ where $\theta$ is the variables pytree, and init is $x, \kappa \mapsto \theta_0$ — symbols: $\theta$ the parameter tree; $\kappa$ the PRNG key; $\theta_0$ the initialised tree. The module class itself contributes only structure: names, shapes-by-inference, and initialiser choices. Internally, linen threads a scope object through the call, collecting param(name, init_fn, shape) declarations into the tree path corresponding to the module path — which is why tree paths mirror code structure (Dense_0/kernel) and why checkpoint formats are stable under refactors that preserve module names.

Collections generalise parameters. The variables dict is keyed by collection: params (trained), batch_stats (BatchNorm running moments — updated by forward passes, not gradients), arbitrary user collections. apply(..., mutable=['batch_stats']) returns (output, updated_collections) — mutation reified as return values, the functional answer to PyTorch's buffers. Similarly, stochastic layers consume named RNG streams (rngs={'dropout': key}), the framework-level form of explicit key threading.

The state problem and NNX

Linen's purity has a usability bill: every stateful thing (params, batch stats, RNG streams, optimizer state) is threaded manually through function signatures — the "train state shuffle". Flax's NNX (2024–) re-introduces Python-object state with graph-based tracking, transform-aware (nnx.jit, nnx.grad), aiming at PyTorch-like ergonomics on JAX semantics; the Flax team recommends NNX for new projects while linen remains fully supported and dominant in existing code (including much of the published JAX research ecosystem — T5X, Scenic lineage). Reading fluency in linen therefore stays mandatory; the mould/dough model transfers to NNX unchanged underneath.

Alternatives map the same design space: Haiku (DeepMind; transform converts stateful-looking code to init/apply pairs) and Equinox (modules are pytrees, parameters as leaves — the minimalist end). All converge on one invariant: whatever the surface syntax, the compiled artefact is a pure function over a parameter pytree.

References

  • Heek et al. (2020–), Flax: a neural network library and ecosystem for JAX.
  • Hennigan et al. (2020), Haiku: Sonnet for JAX — the transform-based alternative.
  • Kidger and Garcia (2021), Equinox: neural networks in JAX via callable PyTrees and filtered transformations.

What to learn next

What to learn next

These follow on from what you just read.

  • JAX and Flax

    Optax

    Optax is JAX's optimizer library — each optimizer is a pair of pure functions plus a state you carry, and complex training recipes snap together from small pieces.

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

  • Natural Language Processing

    What is NLP?

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