JAX and Flax

Random numbers and PRNG keys

JAX has no hidden random state — every random draw takes an explicit key, the same key always gives the same numbers, and fresh keys come from splitting.

Read these first

On this page 5
  1. Why JAX refuses a shared pot
  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.

In JAX, random numbers are not drawn from a shared pot — every draw needs a key you hand over yourself, and the same key always produces the same numbers.

Think of a scratch card. The numbers under the foil were fixed the moment the card was printed. Scratching reveals them; scratching again shows the same ones. Want different numbers? You need a different card.

A JAX key is a scratch card. Ask for random numbers with a key and you get that key's numbers — today, tomorrow, on any machine. New randomness requires splitting: turning one card into two fresh cards, each with its own fixed numbers.

Why JAX refuses a shared pot

NumPy keeps one hidden random state, and every draw quietly advances it. Convenient — and it makes results depend on the order and count of every draw anywhere in your program. One extra draw in a helper function shifts every number after it. Debugging a training run that cannot be replayed is misery; the wider war for reproducibility is fought over exactly this.

JAX also has a deeper problem: its compiler and batching tools rearrange and parallelise your code. A hidden counter that must tick in order would fight every one of those tools. Explicit keys remove the fight — each draw's outcome is pinned by its key, whatever order things run in.

How it works

key = make a key from a seed (say, 42)

draw(key)  ──▶  0.31, -1.2, 0.88        same key ──▶ same numbers, always

split(key) ──▶  key_a , key_b           two fresh, unrelated cards
                 ↓
use key_a for dice roll 1, key_b for dice roll 2

The working rhythm: split, use one half, keep the other for next time. A key you have already used is spent — reusing it repeats its numbers.

A real example you have seen

Multiplayer game replays. When you watch a replay of an online match, every "random" event happens identically — critical hits, item drops alike. The game recorded the seed and replays the same randomness. Bug reports become reproducible; matches become verifiable. Explicit keys give your experiments that replay button.

Remember this

  • A key pins randomness: same key, same numbers, every time.
  • Split a key for fresh randomness; never reuse a spent key.
  • No hidden state means results cannot drift because of unrelated code.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install jax

Outputs verified with jax 0.6.2, CPU. These exact numbers should reproduce for you — that is the entire point of this lesson.

Keys, repetition, and splitting

keys.py
import jax

key = jax.random.key(42)

# Same key, same numbers — every single time, on any machine
print(jax.random.normal(key, (3,)))
print(jax.random.normal(key, (3,)))

# Want fresh numbers? Split the key
key, subkey = jax.random.split(key)
print(jax.random.normal(subkey, (3,)))

# Need many at once? Split into many
keys = jax.random.split(key, 4)
print("4 independent keys, shape:", keys.shape)
Output
[-0.02830462  0.46713185  0.29570296]
[-0.02830462  0.46713185  0.29570296]
[ 0.60576403  0.7990441  -0.908927  ]
4 independent keys, shape: (4,)

The walkthrough

Two identical draws from one key — the scratch-card rule made visible. In NumPy those two lines would differ, because a hidden state advanced between them. In JAX nothing advanced: there is nothing to advance.

key, subkey = jax.random.split(key) is the idiom to burn into your fingers. Split produces two fresh keys; the convention is to use the subkey now and carry the reassigned key forward. The name-shadowing (key on both sides) is deliberate — the old key disappears from reach, so you cannot accidentally reuse it.

split(key, 4) manufactures a batch of independent keys — one per layer to initialise, one per dropout mask, one per vmap lane. Passing a key per lane through vmap gives every batch element its own randomness, deterministically.

jax.random.key(42) (the modern typed form; PRNGKey is the older spelling still seen everywhere) builds the starting card from a seed. One seed at the top of an experiment, splits all the way down: the whole run becomes replayable from one integer.

In training loops, the rhythm appears at every step needing noise — for dropout, this is where each step's mask key comes from:

python
key, dropout_key = jax.random.split(key)

Common mistakes

Reusing a key across draws. Two dropout layers fed the same key drop the same pattern; a data shuffle and a noise draw sharing a key correlate mysteriously. Correlated randomness rarely crashes — it quietly poisons results. Rule: one key, one use, ever.

Forgetting to reassign after split. subkey = jax.random.split(key)[1] while key stays in scope invites reuse. The two-target assignment above is the safe shape.

Seeding like NumPy. There is no jax.random.seed(...) that fixes a global. If you find yourself looking for it, the mental model has not flipped yet: the key is the state, and you carry it.

Expecting NumPy's numbers. Same seed, different generator: JAX (Threefry) and NumPy (PCG64/MT) produce unrelated streams. Cross-library comparisons must compare distributions, not values.

Try it yourself

Draw 10,000 values with jax.random.normal(subkey, (10000,)) and check mean() and std() land near 0 and 1. Then split one key into 1,000 subkeys, draw one number from each, and verify the statistics again — splitting produces streams as clean as consecutive draws.

What to learn next

Researcher — Mathematics and papers.

Counter-based PRNG

JAX's default generator is Threefry-2x32 (Salmon et al. 2011, Parallel random numbers: as easy as 1, 2, 3), a counter-based PRNG: a keyed block cipher applied to a counter, $r_i = E_k(i)$ — symbols: $E_k$ the cipher under key $k$; $i$ the counter/position; $r_i$ the $i$-th random block. Unlike sequential generators (Mersenne Twister, PCG) whose state must advance draw by draw, any position is computable independently in $O(1)$ — random access to the stream. split derives child keys by encrypting under the parent, giving a tree of statistically independent streams without coordination.

This is what buys the properties the beginner block promised: order-independence (a draw's value depends on its key, not on program history), parallel safety (vmap lanes and pmap devices take disjoint keys with zero communication), and reproducibility across execution strategies — the same program jitted, vmapped, or run eagerly yields identical randomness, which a stateful generator cannot promise once operations reorder.

The functional-purity connection

A hidden RNG state is a side effect, and jit compiles pure functions: with global state, either every traced function takes the RNG as hidden input/output (breaking caching and reordering), or randomness silently freezes at trace time — the classic bug where a "random" jitted function returns the same numbers forever. Explicit keys make the dependency visible to the tracer like any other argument. The price is threading keys through every stochastic function — accepted as boilerplate, and absorbed by frameworks: Flax plumbs named RNG streams (params, dropout) through init/apply, so user code supplies one key per stream per call.

Statistical caveats: Threefry's guarantees are empirical (BigCrush-clean) rather than cryptographic at the rounds JAX uses; and naive fold-in patterns (fold_in(key, step)) construct correlated families if used to bypass splitting discipline — fold_in is well-defined and safe for deriving per-step keys, but mixing fold-in and split hierarchies ad hoc deserves thought. The typed key arrays (jax.random.key) exist partly to allow alternative generators (rbg, unsafe_rbg for TPU throughput) behind one API.

References

  • Salmon et al. (2011), Parallel random numbers: as easy as 1, 2, 3, SC'11 — Threefry and the counter-based idea.
  • Claessen and Pałka (2013), Splittable pseudorandom number generators using cryptographic hashing — the splitting formalism.
  • L'Ecuyer and Simard (2007), TestU01 — the BigCrush battery cited for quality claims.

What to learn next