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.
- 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.
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 2The 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
- Control flow: scan, cond and while_loop — carrying keys (and everything else) through compiled loops.
- Random seeds and reproducibility — the same war in PyTorch, fought uphill.
- Building models with Flax — where the key-threading gets absorbed by a framework.
Developer — Code and libraries.
Setup
pip install jaxOutputs 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
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)[-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:
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
- Control flow: scan, cond and while_loop — carrying keys (and everything else) through compiled loops.
- Random seeds and reproducibility — the same war in PyTorch, fought uphill.
- Building models with Flax — where the key-threading gets absorbed by a framework.
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
- Control flow: scan, cond and while_loop — carrying keys (and everything else) through compiled loops.
- Random seeds and reproducibility — the same war in PyTorch, fought uphill.
- Building models with Flax — where the key-threading gets absorbed by a framework.