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.
- 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.
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 placeThe .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
- jit and tracing — the compiler that cashes in on immutability.
- NumPy — the API being mirrored, if slicing and broadcasting feel shaky.
- In-place operations and autograd — what mutation costs the other framework.
Developer — Code and libraries.
Setup
pip install jaxThe CPU wheel is small — nothing like TensorFlow's download. Outputs verified with jax 0.6.2, CPU.
The rule, met head-on
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)[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
- jit and tracing — the compiler that cashes in on immutability.
- NumPy — the API being mirrored, if slicing and broadcasting feel shaky.
- In-place operations and autograd — what mutation costs the other framework.
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
- jit and tracing — the compiler that cashes in on immutability.
- NumPy — the API being mirrored, if slicing and broadcasting feel shaky.
- In-place operations and autograd — what mutation costs the other framework.