Control flow: scan, cond and while_loop
Compiled JAX code cannot branch or loop on values with Python's if and while — lax.cond, lax.scan and lax.while_loop are the versions that trace correctly.
- 8 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.
Python's if and while ask questions about values — but compiled JAX code is recorded before the values exist, so JAX provides its own if and loop that go into the recording.
Think of writing instructions for a courier who will make deliveries next week, alone, with no phone. You cannot write "call me when you reach the gate and I'll decide". Every decision must be written into the sheet now: "IF the gate is locked, leave the parcel with the guard; OTHERWISE ring the bell." The courier carries both branches and picks one on the spot.
lax.cond is that written-in-advance IF. lax.scan and lax.while_loop are written-in-advance loops. They exist because compiled functions are recordings, and a recording cannot pause to ask Python a question.
Why it exists
During tracing, your function runs with stand-in values — shape and type known, actual numbers absent. A Python if x > 0: needs the answer now, at recording time, and the stand-in has none to give. That is the TracerBoolConversionError every JAX learner meets in week one.
The fix is decisions that live inside the recording. cond records both branches and a selector. scan records the loop body once, plus the instruction "run this N times, carrying results forward". It also compiles far faster than a Python loop, whose every turn gets recorded again.
How it works
lax.cond(condition, do_this, else_this, x) both branches recorded,
chosen when values exist
lax.scan(step, start, items) step recorded ONCE:
start ─▶ step ─▶ step ─▶ step ─▶ final carry flows left to right,
│ │ │ per-step outputs collected
▼ ▼ ▼
out 1 out 2 out 3
lax.while_loop(keep_going, step, start) loop until a test says stopThe carry is the parcel passed from one step to the next — a running total, a model state, whatever must survive between turns.
A real example you have seen
A bank passbook being updated. Each transaction takes the previous balance, applies one change, and hands the new balance to the next line. The clerk's rule is written once; the balance is the carry; the printed column of balances is the collected output. That is scan, exactly.
Remember this
- Python
if/whileon values cannot enter compiled JAX code. cond= both branches recorded, one chosen at run time.scan= one recorded step + a carry — the workhorse for loops and sequences.
What to learn next
- Building models with Flax — modules that put scans and conds behind clean layers.
- RNN — the architecture that is a scan wearing a hat.
- jit and tracing — the recording model that made all this necessary.
Developer — Code and libraries.
Setup
pip install jaxOutputs verified with jax 0.6.2, CPU. Deterministic throughout.
The three tools on real shapes of problem
import jax
import jax.numpy as jnp
from jax import lax
# 1. cond: an if that works under jit
def step_fee(amount):
return lax.cond(amount > 1000.0,
lambda a: a * 0.02, # big transfers: 2% fee
lambda a: 10.0, # small ones: flat 10
amount)
print("fee on 500:", jax.jit(step_fee)(500.0))
print("fee on 5000:", jax.jit(step_fee)(5000.0))
# 2. scan: a loop that carries state — here, a running bank balance
def apply_txn(balance, txn):
new_balance = balance + txn
return new_balance, new_balance # (carry, per-step output)
txns = jnp.array([100., -30., -50., 200.])
final, history = lax.scan(apply_txn, 0.0, txns)
print("final balance:", final)
print("balance history:", history)
# 3. while_loop: repeat until a condition breaks
def halve_until_small(x):
return lax.while_loop(lambda v: v > 1.0, lambda v: v / 2.0, x)
print("100 halved until <= 1:", halve_until_small(100.0))fee on 500: 10.0 fee on 5000: 100.0 final balance: 220.0 balance history: [100. 70. 20. 220.] 100 halved until <= 1: 0.78125
The walkthrough
cond(pred, true_fn, false_fn, operand) — both functions are traced, the predicate picks at execution. Under jit, this is what "if" must become. When both sides are cheap element-wise maths, jnp.where(pred, a, b) is lighter — it computes both and selects. Reach for cond when branches are genuinely expensive or structurally different.
scan's contract is strict and worth memorising: the step function takes (carry, one_item) and returns (new_carry, one_output). Feed it the start carry and the stacked items; receive the final carry and the stacked outputs. Check it against the output: running balance 100, 70, 20, 220 — the passbook column.
Why scan over a Python loop? A Python loop under jit unrolls — 10,000 iterations means 10,000 copies of the body in the recording, and compile times measured in coffee breaks. scan records the body once. RNN-style models, optimizer inner loops, and sampling loops in production JAX are scans.
while_loop(cond_fn, body_fn, init) covers loops whose trip count depends on data — nobody knows in advance how many halvings 100 needs. The price of that freedom: no per-step outputs are collected, and reverse-mode gradients cannot flow through it (the recording cannot know how many steps to walk back). scan supports gradients fully; prefer it whenever the length is known.
Common mistakes
Branches with different shapes. cond requires both branches to return the same shape and dtype — 0.0 from one and a (3,) array from the other raises a structure error. Both couriers' instructions must fit the same parcel slot.
Breaking scan's return contract. Returning a bare value instead of the (carry, output) pair produces confusing structure errors. If there is nothing to collect per step, return (new_carry, None).
A changing carry shape. The carry must keep one shape across steps — growing a list, widening an array mid-loop will not trace. Pre-allocate to the final size and fill by index instead.
Reaching for while_loop by default. It is the least capable of the three: no collected outputs, no reverse-mode gradient. Known length → scan. Unknown length but bounded → often still scan to the bound, with a mask.
Try it yourself
Rewrite the balance example so any transaction pushing the balance below zero is skipped — a cond (or jnp.where) inside the scan step. The carry stays a single number; only the step logic changes. Verify with txns = [100., -150., 30.], expecting the −150 to bounce.
What to learn next
- Building models with Flax — modules that put scans and conds behind clean layers.
- RNN — the architecture that is a scan wearing a hat.
- jit and tracing — the recording model that made all this necessary.
Researcher — Mathematics and papers.
Structured control flow as primitives
cond, scan, and while_loop are higher-order primitives in the jaxpr IR: their operands include sub-jaxprs (the traced branches or body), so the recording contains control structure rather than a flattened trace. XLA lowers them to native Conditional and While HLO ops. This is the structured alternative to the two degenerate strategies: full unrolling (Python loops under jit — code size $O(n)$ in trip count, compile time to match) and staying eager (per-op dispatch, no fusion). TensorFlow reached the same destination via AutoGraph's automatic conversion of Python syntax (the tf.function machinery); JAX makes the conversion explicit and manual — one less layer of magic, one more thing to learn.
scan and differentiation
scan(f, c_0, x_{1:n}) computes $c_i, y_i = f(c_{i-1}, x_i)$ — symbols: $c_i$ the carry after step $i$; $x_i$ the $i$-th input slice; $y_i$ the collected output. It is the functional form of an RNN/state-space update, and its reverse-mode derivative is itself a scan run backwards (adjoint state propagation), so gradients cost $O(n)$ time with activation storage $O(n)$ — reducible via jax.checkpoint on the body to $O(\sqrt{n})$-style rematerialisation trade-offs (Chen et al. 2016). Backpropagation through a scan is backpropagation-through-time in the RNN sense, with the same vanishing/exploding-gradient physics along the carry chain.
while_loop is reverse-mode non-differentiable by construction — trip count is data-dependent, so the adjoint's storage cannot be planned. Forward-mode (jvp) works. Fixed-point tricks (implicit function theorem — Blondel et al. 2022) recover gradients for convergent while-loops without unrolling, the standard workaround in equilibrium models.
Performance geometry
cond executes one branch; jnp.where-style select executes both. On accelerators, select frequently wins for cheap branches by avoiding divergent control flow, while cond wins when a branch contains a matmul you would rather skip. Under vmap, cond over a batched predicate is lowered to select semantics anyway (both branches run for the batch) — a documented, occasionally surprising cost cliff: per-example "skipping" cannot exist in lockstep SIMD execution.
References
- Chen et al. (2016), Training deep nets with sublinear memory cost — checkpointing through long scans.
- Blondel et al. (2022), Efficient and modular implicit differentiation.
- Gu et al. (2022), Efficiently modeling long sequences with structured state spaces — scan as the compute pattern of modern sequence models.
What to learn next
- Building models with Flax — modules that put scans and conds behind clean layers.
- RNN — the architecture that is a scan wearing a hat.
- jit and tracing — the recording model that made all this necessary.