JAX and Flax

JAX arrays and immutability

JAX arrays look like NumPy arrays with one rule changed — you can never edit one in place, only produce a corrected copy.

Read these first

On this page 5
  1. Why JAX chose this
  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 array is like a NumPy array with one rule changed: once created, it can never be edited — you can only make a new array with the change applied.

Think of a bank passbook. When a transaction is wrong, the bank never erases the ink and writes over it. It adds a fresh correcting entry, and the old line stays visible forever. Records that nobody can overwrite are called immutable — unchangeable after creation.

JAX arrays are passbook entries. NumPy arrays are a pencil ledger — rub out, write over. Same numbers, opposite philosophy.

Why JAX chose this

JAX's whole trick is transforming your functions: making them fast, batching them, differentiating them — the transformations fill the rest of this section. All of them work by reading your calculation as a clean sequence of steps.

In-place edits poison that reading. If a value can change under your feet, the order of every step suddenly matters in hidden ways, and reorganising the steps for speed becomes unsafe. Forbid editing, and every step's inputs and outputs are explicit — the calculation becomes safe to rearrange, split across devices, and replay.

The cost you would expect — endless copying — mostly never happens. JAX watches the whole calculation and skips copies nobody could notice.

How it works

NumPy:   scores[0] = 99        the array itself is changed
                               (anyone holding it sees 99)

JAX:     new = scores.at[0].set(99)
              ─────────────────────
         scores  → unchanged, still the old numbers
         new     → a fresh array with 99 in place

The .at[...] form is the correcting entry: say what you want changed, receive a new passbook page, keep or discard the old one.

A real example you have seen

Ever edited a shared document where "track changes" is on? Nothing is truly deleted — every edit becomes a visible addition, and you can rewind to any earlier state. Immutable data gives programs that same rewind-and-replay safety, and JAX spends that safety on speed.

Remember this

  • JAX arrays are immutable — created once, never edited.
  • Updates go through .at[index].set(value), which returns a new array.
  • The no-editing rule is what makes JAX's speed tricks safe to apply.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install jax

The CPU wheel is small — nothing like TensorFlow's download. Outputs verified with jax 0.6.2, CPU.

The rule, met head-on

immutability.py
import jax.numpy as jnp

x = jnp.array([10., 20., 30.])
print(x, x.dtype)

try:
    x[0] = 99.0
except TypeError as err:
    print("error:", str(err)[:68])

y = x.at[0].set(99.0)          # returns a NEW array
print("y:", y)
print("x untouched:", x)

z = x.at[1].add(5.0)           # add, mul, min, max all exist
print("z:", z)
Output
[10. 20. 30.] float32
error: JAX arrays are immutable and do not support in-place item assignment
y: [99. 20. 30.]
x untouched: [10. 20. 30.]
z: [10. 25. 30.]

The walkthrough

jax.numpy mirrors NumPy on purpose. Most NumPy code ports by changing the import — jnp.dot, jnp.mean, broadcasting, slicing all behave as your NumPy reflexes expect. Immutability is the headline exception, so it gets this whole lesson.

The error is a design decision, not a limitation. JAX could have supported x[0] = 99 and chose to refuse. The message even suggests the fix: use .at.

x.at[0].set(99.0) reads as: at index 0, set 99, give me the result. Any indexing you can read with can write with — slices x.at[1:].set(0), boolean masks, fancy indexing. The companions add, mul, min, max fuse read-modify-write into one step.

The copy usually is not real. Inside compiled code, JAX proves when the old array is never used again, and then performs the update in place behind the scenes. You write as if copying; the machine edits. This "donation" of dead buffers is why functional style does not mean slow style — the proving happens under jit.

Two smaller surprises worth meeting today: default floating dtype is float32 (not NumPy's float64) — doubles need an opt-in flag. And operations run asynchronously: a jnp expression returns before the maths finishes, which is why timing code needs .block_until_ready(). Both bite in later lessons; recognise them now.

Common mistakes

Porting NumPy code and hitting the TypeError. Every arr[i] = v, arr += ... on a slice, or np.fill pattern must become .at. The error message is loud, so this one at least fails fast.

Discarding the result. x.at[0].set(99) alone changes nothing — the new array must be caught: x = x.at[0].set(99). Sequential updates chain the same way. The silent no-op version is the dangerous cousin of the loud TypeError.

Writing element-by-element loops. A Python loop of single .at updates creates one array per step — correct, and slow, exactly like element-wise loops in NumPy. Think in whole-array operations first; reach for .at when you truly need targeted surgery.

Expecting float64. jnp.array([1.0]) is float32. Comparisons against NumPy results then differ around the 7th digit. For genuine float64 work: jax.config.update("jax_enable_x64", True) at program start.

Try it yourself

Take scores = jnp.arange(10.0) and produce a version where every value above 5 is capped at 5 — once with a boolean mask through .at, once with jnp.minimum(scores, 5). Confirm both leave scores untouched, and decide which reads better.

What to learn next

Researcher — Mathematics and papers.

Purity as the enabling contract

JAX's transformations (jit, grad, vmap, pmap) are defined on pure functions: outputs depend only on inputs, no observable mutation, no side effects. Immutable arrays enforce the data half of that contract. The payoff is referential transparency — any subexpression may be substituted by its value — which licenses the rewrites XLA performs: common-subexpression elimination, fusion, reordering across the program, dead-code elimination. Under mutation, each of these requires alias analysis to prove safety; under immutability they are sound by construction (the classical argument from functional programming — Backus 1978, Can programming be liberated from the von Neumann style?).

The update cost model

Semantically, x.at[i].set(v) is functional update: $y = x[0{:}i] \mathbin{\Vert} v \mathbin{\Vert} x[i{+}1{:}]$, an $O(n)$ copy for an array of $n$ elements. Operationally, XLA's buffer assignment performs liveness analysis; when $x$ is dead after the update (the common case in a compiled function), the copy is elided and the write happens in place — $O(1)$ amortised for the mutation pattern of a training loop. Outside jit, in op-by-op ("eager") execution, the copy is generally real: micro-benchmarks of .at in a Python loop measure the semantics, not the compiled behaviour — a standard benchmarking trap.

Explicit control exists at function boundaries: jax.jit(f, donate_argnums=0) donates the input buffer, telling XLA the caller abandons it, enabling in-place reuse for arguments too (an error to reuse the donated array afterwards).

Contrast with the mutation-tracking alternative

PyTorch permits in-place ops and pays with machinery: version counters on tensors, runtime errors when autograd detects an overwritten value needed for backward (in-place operations and autograd), and semantic limits on torch.compile graph capture around aliasing. JAX moves the same complexity from runtime checking to a compile-time guarantee. The trade: JAX forbids a convenient syntax; PyTorch forbids, at runtime and unpredictably, a subset of its uses.

Asynchronous dispatch note for measurement: JAX enqueues work and returns futures-like arrays; wall-clock measurements must fence with .block_until_ready() or they measure enqueue time (documented in the JAX benchmarking guide; the same discipline CUDA users know from torch.cuda.synchronize).

References

  • Frostig, Johnson, Leary (2018), Compiling machine learning programs via high-level tracing — the JAX design paper.
  • Sabne (2020), XLA: compiling machine learning for peak performance — the buffer assignment and fusion layer.
  • Backus (1978), Can programming be liberated from the von Neumann style? — the founding argument for purity.

What to learn next

What to learn next

These follow on from what you just read.

  • JAX and Flax

    jit and tracing

    jax.jit runs your function once with stand-in values, records every operation, and compiles the recording into fast machine code it replays from then on.

  • JAX and Flax

    grad and value_and_grad

    jax.grad takes a function and returns a new function — its derivative — and value_and_grad returns both answer and slope in one pass.

  • JAX and Flax

    vmap: write the code for one example

    vmap takes a function written for a single example and returns one that handles a whole batch — no loops, no reshaping, full speed.