JAX#
JAX is our tool of choice for high-performance numerical work in QuantEcon lectures.
This page covers three things:
Outline — what JAX is, when a lecture should use it, and how to set one up.
JAX Style Guide — how to write clean, idiomatic, modern JAX.
Converting from Numba — how to migrate existing Numba code.
Outline#
What JAX is#
JAX provides a NumPy-compatible array API (jax.numpy) plus a set of
function transformations:
jax.jit— compile a function to fused machine code via XLA, for CPU or GPUjax.vmap— vectorize a function over a new array axis, without materializing intermediate arraysjax.grad— automatic differentiation
Alongside these there is jax.lax, which is not a transformation but a module
of primitives that the compiler understands — most importantly the control
flow operations scan, fori_loop and while_loop, which we use in place of
Python loops.
The price of these transformations is that JAX requires pure functions and immutable arrays. Code must be written in a functional style, and this style is not always the most readable one.
That trade-off is the whole basis of the guidance below.
Important
While performance is a goal in switching to JAX, our focus should be on writing code that is clean, simple to understand, highly readable, and as close to the math as possible.
This may mean not always implementing the most performant JAX code available, particularly if a section of code does not require performance — such as generating plots.
When to use JAX#
Use JAX when at least one of the following holds.
There is a real bottleneck. The lecture contains a computation whose runtime is actually a constraint — a value function iteration, a large simulation, a big grid search. Not a 100 × 100 plotting grid.
The work vectorizes or parallelizes. The bottleneck is array-shaped, or consists of many independent problems that
vmapcan batch. JAX’s advantage comes from fusion and parallelism, so a long chain of small dependent steps has nothing for it to exploit — and on a GPU such a loop is actively slower than the Numba equivalent.You need automatic differentiation. Gradients, Jacobians, sensitivities, or anything built on them (maximum likelihood, neural networks, gradient-based solvers).
The lecture is teaching JAX itself, or is a stepping stone towards a later high-performance lecture.
If none of these hold, do not convert the lecture. NumPy — or NumPy plus Numba — is the right answer, and the resulting code will be shorter and easier to read.
Note
Differentiability is also a forward-looking reason to lean toward JAX at the margin. Autodiff is exact, cheap, and composes cleanly to high dimensions and higher orders, and we expect to exploit it more and more — for sensitivity analysis, calibration, and estimation — as we grow used to its power. So when a lecture has a differentiable object and the decision to convert is otherwise balanced, that trajectory is a reasonable tie-breaker: a JAX lecture is a ready foundation for it, whereas a NumPy one would need a rewrite. This tips genuinely marginal cases only — it does not override the tests above.
Warning
Substituting jnp for np is not converting a lecture to JAX.
If the numerical work still runs through SciPy, scipy.stats, or NumPy-based
QuantEcon.py classes such as Kalman and LinearStateSpace, then the arrays are
merely being shuttled between host and device at every call. The lecture gets
JAX’s constraints and none of its benefits.
A genuine conversion moves the computation — not just the array constructors.
When not to use JAX#
Some concrete anti-patterns, all of which we have seen in review:
Rewriting a library call by hand. scipy.stats.multivariate_normal(μ, Σ).pdf(pos)
is one line and obviously the mathematics. Replacing it with a hand-rolled density
using vmap, jnp.linalg.inv and jnp.linalg.det costs thirty lines and buys
nothing when the grid has ten thousand points.
Functional updates inside a Python loop. This is the worst of both worlds:
# ❌ Slow and unidiomatic: outside jit, each .at[].set() copies the whole array
z = jnp.empty(T)
for t in range(T):
z = z.at[t].set(compute_something(t))
Outside jit there is no compiler to elide the copy, so this is \(O(T^2)\) work
where NumPy is \(O(T)\). If the loop body genuinely must run in Python — because it
calls SciPy, or steps a stateful object — use NumPy:
# ✅ Clear and fast
z = np.empty(T)
for t in range(T):
z[t] = compute_something(t)
Converting only the plots. Figure generation is never the bottleneck. Leave plotting code in NumPy and Matplotlib where it is clearest.
Adding the GPU admonition to a lecture that has no GPU work. It misleads readers and implies CI configuration the lecture does not need.
Vectorized versus sequential work#
The case for JAX is much stronger for vectorized operations than for sequential ones. It is worth seeing why.
Vectorized: JAX wins#
Consider maximizing
over a fine grid on \([-3, 3]^2\).
The NumPy approach uses meshgrid:
grid = np.linspace(-3, 3, 3_000)
x_mesh, y_mesh = np.meshgrid(grid, grid)
z_max = np.max(f(x_mesh, y_mesh))
This is correct but memory-hungry: the two mesh grids are two-dimensional, and
NumPy’s eager execution creates further intermediate arrays of the same size at
each step of f.
Numba avoids the memory cost with a plain loop, and parallelizes with prange:
@numba.jit(parallel=True)
def compute_max(grid):
n = len(grid)
m = -np.inf
for i in numba.prange(n):
for j in range(n):
x, y = grid[i], grid[j]
m = max(m, np.cos(x**2 + y**2) / (1 + x**2 + y**2))
return m
JAX gets both the memory efficiency and the parallelism, using vmap in place
of meshgrid:
@jax.jit
def compute_max(grid):
# Max over all x, for a given y
compute_column_max = lambda y: jnp.max(f(grid, y))
# Vectorize so we can call on all y simultaneously
column_maxes = jax.vmap(compute_column_max)(grid)
return jnp.max(column_maxes)
No two-dimensional array is ever created, and because everything sits under a
single jax.jit, XLA fuses the whole calculation into one kernel.
JAX is the winner here. It dominates NumPy on both speed and memory, and dominates Numba once a GPU is available.
Note
Numba can target GPUs through numba.cuda, but then we must parallelize by hand.
For most problems in economics, econometrics and finance it is far better to hand
the parallelization to the JAX compiler than to hand-code it ourselves.
Sequential: Numba wins#
Now consider iterating the quadratic map \(x_{t+1} = \alpha x_t (1 - x_t)\) and storing the whole path. Here is the Numba version:
@numba.jit
def qm(x0, n, α=4.0):
x = np.empty(n+1)
x[0] = x0
for t in range(n):
x[t+1] = α * x[t] * (1 - x[t])
return x
This is exactly how most people think about the algorithm, and it is fast.
JAX can match it — but only on the CPU, and only if we say so explicitly:
cpu = jax.devices("cpu")[0]
# Pin the input to the CPU, which keeps the whole computation there
x0_cpu = jax.device_put(0.1, cpu)
@partial(jax.jit, static_argnames=("n",))
def qm_jax(x0, n, α=4.0):
def update(t, x):
return x.at[t + 1].set(α * x[t] * (1 - x[t]))
x = jnp.empty(n + 1).at[0].set(x0)
return lax.fori_loop(0, n, update, x)
x = qm_jax(x0_cpu, n)
This fori_loop reads much like the Numba version — fill a preallocated array
one element at a time — which is why we prefer it to lax.scan here. The
x.at[t + 1].set(...) update has functional semantics, but because it runs
inside jit the compiler performs it in place, with no per-step copy.
Warning
Left on a GPU, this loop is very slow.
The calculation is a long chain of tiny dependent operations — each step needs the previous one, so there is nothing to parallelize, and every step pays kernel launch overhead on a device built for wide parallel work.
The jax.device_put call above is not incidental: without it, the JAX version
loses badly to Numba on a GPU-equipped build machine.
So even in the best case — on the CPU, with the device pinned by hand — JAX
only draws level with Numba here. And the code is still harder to read than the
plain Numba loop: n must be held static because it determines the array size,
the array is filled through functional x.at[t + 1].set(...) updates rather than
ordinary assignment, and the reader now has to understand device placement too.
Guidance
For inherently sequential work that does not need autodiff, prefer Numba.
At best JAX ties on speed — and only after being pinned to the CPU — while the syntax is considerably less clear.
Reach for JAX when the work vectorizes, when a GPU can actually be exploited, or when you need gradients through the sequence.
See also
The full comparison, with live timings on the build machine, is in the NumPy vs Numba vs JAX lecture.
Setting up a JAX lecture#
Do not pip install jax at the top of the lecture. That may install
jax[cpu], which runs but is not the configuration we build with. The
jax[gpu] package is installed for us via Docker and GitHub Actions.
Instead, include the shared GPU admonition, which lives in each lecture repo at
lectures/_admonition/gpu.md:
```{include} _admonition/gpu.md
```
Always use the include. Never paste the admonition text into a lecture: a
copy stops tracking the shared file the moment the wording changes, and we then
have as many versions of it as there are lectures.
For the same reason the text is deliberately not reproduced here — read
lectures/_admonition/gpu.md in the relevant repo, or look at any lecture that
already includes it, to see how it renders.
See also
Make sure the repo is compliant with the following configuration requirements such as Support files
Please consult with Matt McKay should you need to update these settings.
JAX Style Guide#
This section covers the house conventions. It is deliberately short — it is not a substitute for the JAX documentation.
All examples follow QuantEcon’s variable naming conventions, in particular the use of Unicode Greek letters (α, β, γ, …).
Imports and precision#
from functools import partial
from typing import NamedTuple
import jax
import jax.numpy as jnp
from jax import lax
Note
JAX defaults to 32 bit floats, unlike NumPy.
Single precision is fine for a good deal of iterative work, including much value function iteration, and it is cheaper — 32 bit arrays are faster and use half the memory, and the gap is largest on a GPU. So do not reach for 64 bit reflexively.
It becomes marginal when error accumulates over very many iterations, or when a
fixed point is approached slowly — as when a discount factor sits close to one.
Note also that float32 carries roughly seven significant digits, so a convergence
tolerance much below 1e-6 cannot be met at all.
Test it rather than guessing. Run the computation both ways and check that the answers agree to the accuracy the lecture actually claims:
jax.config.update("jax_enable_x64", True)
This must be set at the top of the lecture, before any array is created.
If 64 bit turns out to be necessary, it is worth a sentence in the text saying so — it is useful for readers to see where precision starts to bite.
Model structure: NamedTuple plus a factory function#
Store model primitives in a NamedTuple and operate on them with free functions.
A NamedTuple is immutable, is a valid JAX pytree, and can be passed straight
through jit boundaries — so it replaces both jitclass and ordinary classes.
class Model(NamedTuple):
α: float
β: float
grid: jnp.ndarray
def create_model(α=0.5, β=0.95, grid_size=100):
"""Build a model instance, validating parameters."""
if not 0 < β < 1:
raise ValueError("β must be in (0, 1)")
grid = jnp.linspace(0.1, 10.0, grid_size)
return Model(α=α, β=β, grid=grid)
@jax.jit
def update(model, state):
return model.β * state ** model.α
Keep validation in the factory function, not in the jitted functions — Python
if statements on traced values will fail under jit.
Pure functions and functional updates#
Functions must not mutate their inputs or depend on global mutable state.
# ❌ Mutating an input
def bad_update(state, shock):
state[0] += shock
return state
# ✅ Returning new data
def good_update(state, shock):
return state.at[0].add(shock)
Note
Although x.at[i].set(v) looks like it allocates a new array every time, inside a
jit-compiled function XLA recognises that the old array is dead and performs the
update in place.
That optimization is only available under jit. Outside it, .at[].set() copies —
see the anti-pattern above.
jit and static arguments#
Decorate the outermost function you want compiled, not every small helper — XLA will inline and fuse the rest.
Arguments that determine array shapes or control flow must be static, so that JAX specializes the compiled code on their values:
@partial(jax.jit, static_argnames=("n",))
def simulate(x0, n, α=4.0):
...
Every distinct value of a static argument triggers a fresh compilation, so keep the set of static arguments small.
vmap instead of meshgrid#
Prefer jax.vmap to broadcasting over a materialized mesh grid. It expresses the
intent more directly and avoids allocating two-dimensional arrays:
# ❌ Allocates two (n, n) grids plus intermediates
x_mesh, y_mesh = jnp.meshgrid(grid, grid)
z = jnp.max(f(x_mesh, y_mesh))
# ✅ Only the flat grid is ever allocated
z = jnp.max(jax.vmap(lambda y: jnp.max(f(grid, y)))(grid))
Scalar kernel plus nested vmap#
This is our standard way of writing a Bellman operator, and more generally any calculation that would otherwise be a nest of loops over states and actions.
Write the kernel for a single state and a single action, so that it reads like the equation it implements:
def _B(v, model, i, j):
"""
The right-hand side of the Bellman equation at one state, before
maximization. Indices (i, j) -> (state, action).
"""
...
return value
Then apply jax.vmap once per loop. In in_axes, a 0 marks the argument
being mapped over and None holds an argument fixed. Map the innermost loop
first:
# v model i j
_B_j = jax.vmap(_B, in_axes=(None, None, None, 0))
_B_ij = jax.vmap(_B_j, in_axes=(None, None, 0, None))
@jax.jit
def B(v, model):
"Value of every action at every state; shape (n_states, n_actions)."
return _B_ij(v, model, jnp.arange(n_states), jnp.arange(n_actions))
Now reduce. The Bellman operator and the greedy policy are the maximum and the maximizer of the same array:
@jax.jit
def T(v, model):
return jnp.max(B(v, model), axis=-1)
@jax.jit
def get_greedy(v, model):
return jnp.argmax(B(v, model), axis=-1)
Why this rather than one large broadcast expression: the reader of _B never
has to decode [:, None] / [None, :] placement to recover the mathematics,
each vmap line corresponds visibly to one loop, and T and get_greedy stop
being near-duplicates of each other — a very common source of drift.
Tip
Map over indices when the kernel indexes into v (e.g. v[i, :]), and over
grid values when it interpolates (e.g. jnp.interp(z, x_grid, v)). The
latter removes a layer of indirection, so prefer it where it applies.
Use jnp.where(feasible, value, -jnp.inf) for constraints, so infeasible
actions are never selected by the maximization but the array stays rectangular.
Control flow#
Python loops inside a jitted function are unrolled at trace time, which makes
compilation slow (or impossible, when the trip count is dynamic). Use lax:
Construct |
Use when |
|---|---|
|
A fixed number of steps — the default choice. It reads like an ordinary |
|
A fixed number of steps, when you prefer the functional form that stacks the per-step outputs for you |
|
The number of steps depends on the data (e.g. convergence) |
Keep the carry type and shape constant across iterations.
Note
scan and fori_loop both lower to the same underlying loop primitive under
XLA, so for a reduction (a loop that returns only its final carry) there is
typically no performance difference between them — the choice is one of
readability, and fori_loop usually wins because it is the smaller step from an
ordinary for loop. The x.at[t].set(...) update used to build a full path has
functional semantics, but inside jit XLA compiles it to an in-place write, so
it allocates no fresh array per step (outside jit it would — see the
anti-pattern above).
Warning
Always give lax.while_loop an iteration bound as well as a convergence
criterion. A while_loop that never satisfies its exit condition will hang
indefinitely, and inside jit there is no way to interrupt it.
@jax.jit
def compute_fixed_point(model, x_init, tol=1e-6, max_iter=10_000):
"""Iterate the update rule to convergence, or until max_iter."""
def cond(state):
i, x, error = state
return (error > tol) & (i < max_iter)
def body(state):
i, x, error = state
x_new = update(model, x)
return i + 1, x_new, jnp.max(jnp.abs(x_new - x))
i, x, error = lax.while_loop(cond, body, (0, x_init, tol + 1.0))
return x
Tip
Note that update is captured from the enclosing scope rather than passed in as
an argument.
Do not write general-purpose helpers that take a function as an argument to a jitted routine. Functions are not valid JAX types, so they have to be marked static, which triggers a fresh compilation for every callable passed in and makes the helper fragile. Close over the update rule, or write a small concrete driver for each model.
Device placement#
By default JAX puts arrays on the fastest available device — the GPU, on our build machines. That is what we want for vectorized work, and the wrong choice for a long sequential loop, where per-step launch overhead dominates and there is no parallelism to recover it.
If a lecture contains a sequential computation that must stay in JAX, pin it to the CPU and explain why:
cpu = jax.devices("cpu")[0]
x0 = jax.device_put(0.1, cpu) # keeps the whole computation on the CPU
Randomness#
JAX has no global random state — every draw takes an explicit key.
Use jax.random.key, not the older jax.random.PRNGKey. The two are
interchangeable in practice, but key produces a typed key array and is the
current recommended interface.
key = jax.random.key(1234)
shocks = jax.random.normal(key, (100,))
Never reuse a key for two different draws — the results will be identical. Split instead:
key, subkey = jax.random.split(key)
z = jax.random.normal(subkey, (100,))
To draw many independent batches, split once into an array of keys and vmap:
keys = jax.random.split(key, num=1000)
paths = jax.vmap(simulate_one_path)(keys)
This mirrors the direction NumPy has taken with np.random.default_rng() — see
NumPy random number generation.
Timing#
JAX dispatches asynchronously, so a timer will record dispatch time rather than
execution time unless you block. Always call block_until_ready(), and always
time a second run so that compilation is excluded:
with qe.Timer():
# First run --- includes compile time
x = compute(data).block_until_ready()
with qe.Timer():
# Second run --- execution only
x = compute(data).block_until_ready()
Converting from Numba#
Decide first#
Before converting anything, check the lecture against the when to use JAX checklist. Numba code that is fast, readable, and sequential should generally stay as it is.
The motivation for converting is performance-critical, vectorizable work, or a need for autodiff — not uniformity for its own sake.
jitclass becomes NamedTuple#
This is usually the largest win. jitclass requires a type specification, ties
data to methods, and is awkward to compose.
# ❌ Numba
from numba import jitclass, float64
spec = [('α', float64), ('β', float64)]
@jitclass(spec)
class Model:
def __init__(self, α, β):
self.α, self.β = α, β
def update(self, x):
return self.β * x ** self.α
# ✅ JAX
class Model(NamedTuple):
α: float
β: float
@jax.jit
def update(model, x):
return model.β * x ** model.α
Loops become array operations or lax#
In order of preference:
Eliminate the loop. Most Numba loops over a grid are really array expressions. Write them as such.
Use
vmapwhen the loop body is a self-contained function of the loop index.Use
lax.scan/fori_loop/while_loopwhen the iteration is genuinely sequential.
# ❌ Numba
@numba.jit
def compute_value(α, β, data):
result = 0.0
for i in range(len(data)):
result += α * data[i] + β
return result
# ✅ JAX --- the loop was never necessary
@jax.jit
def compute_value(α, β, data):
return jnp.sum(α * data + β)
Array mutation becomes functional update#
arr[0] = 5 # ❌ Mutation
arr = arr.at[0].set(5) # ✅ Functional update
arr += 1 # ❌ In place
arr = arr + 1 # ✅ Pure
Random draws become explicit keys#
# ❌ Numba / NumPy
rng = np.random.default_rng(1234)
shocks = rng.standard_normal(100)
# ✅ JAX
key = jax.random.key(1234)
shocks = jax.random.normal(key, (100,))
Worked example#
A simple asset pricing fixed point, converted from jitclass to JAX.
Before (Numba):
from numba import jitclass, float64
import numpy as np
spec = [('β', float64), ('α', float64), ('prices', float64[:])]
@jitclass(spec)
class AssetPricingModel:
def __init__(self, β=0.95, α=0.5):
self.β, self.α = β, α
self.prices = np.array([0.0])
def solve_prices(self, dividends, tol=1e-6):
prices = np.zeros_like(dividends)
for i in range(1000):
new_prices = self.β * (dividends + self.α * prices)
if np.max(np.abs(new_prices - prices)) < tol:
break
prices[:] = new_prices
self.prices = prices
return prices
After (JAX):
class AssetPricingModel(NamedTuple):
β: float = 0.95
α: float = 0.5
def create_asset_model(β=0.95, α=0.5):
if not 0 < β < 1:
raise ValueError("β must be in (0, 1)")
return AssetPricingModel(β=β, α=α)
@jax.jit
def solve_prices(model, dividends, tol=1e-6, max_iter=1_000):
"""Solve for asset prices by fixed point iteration."""
def cond(state):
i, prices, error = state
return (error > tol) & (i < max_iter)
def body(state):
i, prices, error = state
new_prices = model.β * (dividends + model.α * prices)
return i + 1, new_prices, jnp.max(jnp.abs(new_prices - prices))
initial_state = (0, jnp.zeros_like(dividends), tol + 1.0)
i, prices, error = lax.while_loop(cond, body, initial_state)
return prices
model = create_asset_model(β=0.95, α=0.5)
dividends = jnp.array([1.0, 1.1, 0.9, 1.2])
prices = solve_prices(model, dividends)
Note that the iteration count is bounded, so the loop cannot hang, and that validation lives in the factory function rather than inside the jitted routine.
Review checklist#
Before opening a conversion PR, check that:
The lecture satisfies at least one item on the when to use JAX checklist.
The computation moved to JAX — not just the array constructors. There are no residual SciPy or NumPy-based QuantEcon.py calls sitting in the hot path.
No
.at[].set()appears inside a Python loop.Plotting and other non-critical code was left alone.
Precision was tested, not assumed: if
jax_enable_x64is set, the results genuinely needed it; if it is not, 32 bit was checked against 64 bit.Every
lax.while_loophas an iteration bound.Any sequential loop that has to stay in JAX is pinned to the CPU with
jax.device_put, and the lecture says why.The GPU admonition is included via
{include} _admonition/gpu.md, and the lecture actually benefits from a GPU.Figures in the rendered preview match the live site in shape and magnitude.