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.
- 7 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.
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.
initcreates the parameter tree;applyruns the model with it.- Parameters are a plain pytree — everything from the previous lessons applies to them.
What to learn next
- Optax — updating these parameter trees properly.
- A full training loop in JAX — module, optimizer and loop assembled.
- Weight initialisation — why init's random numbers are chosen with care.
Developer — Code and libraries.
Setup
pip install flaxThis 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
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))){'params': {'Dense_0': {'bias': (8,), 'kernel': (3, 8)}, 'Dense_1': {'bias': (1,), 'kernel': (8, 1)}}}
output shape: (4, 1)
param count: 41The 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:
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
- Optax — updating these parameter trees properly.
- A full training loop in JAX — module, optimizer and loop assembled.
- Weight initialisation — why init's random numbers are chosen with care.
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
- Optax — updating these parameter trees properly.
- A full training loop in JAX — module, optimizer and loop assembled.
- Weight initialisation — why init's random numbers are chosen with care.