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.

On this page 5
  1. Why it exists
  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.

vmap takes a function you wrote for one single example and hands back a function that processes a whole stack of examples at once.

Think of a teacher who wrote careful instructions for grading one answer sheet: check question 1, award marks, total them. Exam season arrives with 300 sheets. She does not rewrite the instructions for stacks. She hands the single-sheet instructions to a grading team that applies them across the pile in one go.

vmap is the grading team. You write the one-example version, which is the straightforward version to write and to check. vmap turns it into the batch version.

Why it exists

Neural networks process examples in batches for speed. So everyone's code ends up full of a batch dimension — an extra first direction on every array counting "which example". Writing code where every operation carries this extra direction is a constant tax: harder to write, harder to read, and the source of endless shape bugs.

The loop alternative — "for each example, run the function" — reads beautifully and runs terribly, because loops in Python are slow. vmap gives you both halves: single-example clarity, batched speed. Under the hood it rewrites your operations to act across the pile — it never loops.

How it works

you write:      distance(one_point, one_warehouse)          easy to think about

vmap gives:     distances(many_points, one_warehouse)
                 in_axes says which inputs are piles:
                 (0, None)  =  first input is a pile along axis 0,
                               second is shared by everyone

The in_axes setting answers one question per argument: is this a pile of different values, or one value everyone shares?

A real example you have seen

A food delivery app computes your distance to every nearby restaurant the moment you open it. Somebody wrote "distance between one customer and one restaurant" — the formula from school. The batched machinery applies it across hundreds of restaurants in a blink. Nobody wrote a distances-to-many formula.

Remember this

  • Write for one example; vmap makes it work for a pile.
  • in_axes marks each argument: batched along an axis, or shared (None).
  • It rewrites operations across the batch — it is not a hidden loop.

What to learn next

Developer — Code and libraries.

Setup

bash
pip install jax

Outputs verified with jax 0.6.2, CPU. Deterministic throughout.

One-example code, batched two ways

vmap_distances.py
import jax
import jax.numpy as jnp

# Distance from one delivery point to one warehouse — ONE example each
def distance(point, warehouse):
    return jnp.sqrt(jnp.sum((point - warehouse) ** 2))


points = jnp.array([[0., 0.], [3., 4.], [6., 8.]])     # 3 delivery points
warehouse = jnp.array([0., 0.])

# Batch over points (axis 0), keep warehouse fixed (None)
batched = jax.vmap(distance, in_axes=(0, None))
print("distances:", batched(points, warehouse))

# Batch over BOTH: pairwise table via two nested vmaps
warehouses = jnp.array([[0., 0.], [10., 0.]])
pairwise = jax.vmap(jax.vmap(distance, in_axes=(None, 0)), in_axes=(0, None))
print("3x2 distance table:")
print(pairwise(points, warehouses))
Output
distances: [ 0.  5. 10.]
3x2 distance table:
[[ 0.       10.      ]
 [ 5.        8.062258]
 [10.        8.944272]]

The walkthrough

distance knows nothing about batches. It takes one point, one warehouse, returns one number — the version you can check with the 3-4-5 triangle in your head (and there it is: 5.0 in the output).

in_axes=(0, None) reads argument by argument: points is a pile, batched along axis 0; warehouse is shared. The output gains the batch axis: three numbers. Shape story: (3, 2), (2,) → (3,).

Nested vmap builds the cross product. The inner vmap(distance, in_axes=(None, 0)) maps one point over all warehouses; the outer maps that over all points. Result (3, 2): every point against every warehouse, written with zero loops and zero broadcasting acrobatics. Compare the manual way — inserting fake axes so broadcasting does the work: possible, and the reshaping is exactly the error-prone part vmap removes.

Composition is the superpower. jax.vmap(jax.grad(loss_fn)) gives per-example gradients — the gradient machinery applied example-wise, in one vectorised pass. jit(vmap(f)) compiles the batched version. Transformations stack in any order that makes sense.

out_axes (default 0) controls where the batch axis lands in the output, for when downstream code expects it elsewhere.

Common mistakes

Wrong in_axes for an argument that is genuinely shared. Marking warehouse as 0 here batches over its coordinates — shapes may even align by luck, producing wrong numbers silently. For each argument ask: does each example bring its own, or is it shared? Own → axis; shared → None.

Batching along the wrong axis. Data stored examples-in-columns needs in_axes=1, not 0. If results look transposed or shapes fail, print shapes and find where the example axis truly lives.

Using vmap where broadcasting already works. prices * 0.9 needs no vmap — element-wise operations batch themselves. vmap earns its keep when the per-example function is genuinely non-trivial: an index lookup, a grad, a small model.

Hiding a Python loop inside. If the mapped function itself loops over examples or mutates state, vmap cannot vectorise the intent. The function must treat its input as one example, purely.

Try it yourself

Write nearest(point, warehouses) returning jnp.argmin of distances from one point — using the inner vmap from above. Then vmap it over points to label each delivery with its closest warehouse. Two lines of new code; check the answer against the printed table.

What to learn next

Researcher — Mathematics and papers.

The batching rule, formally

vmap is a program transformation over the traced jaxpr: each primitive $p$ carries a batching rule mapping the pair (inputs with batch dimensions, dimension indices) to (batched output, output dimension). Tracing propagates a BatchTracer carrying (value, batch_dim) through the function; primitive by primitive, the rule rewrites, e.g., a dot product into a batched matmul, an index-gather into a batched gather. The semantics are exactly:

$$ \text{vmap}(f)(x_{1:B}) = \big(f(x_1), \dots, f(x_B)\big) $$

with $B$ the batch size and $x_i$ the $i$-th slice along the mapped axis — but the implementation is a single program over arrays with one extra dimension, not $B$ calls. Cost is therefore that of the equivalent hand-batched program: one fused kernel launch per primitive, memory $O(B)$ times the single-example footprint. A Python loop, by contrast, pays dispatch overhead $B$ times and defeats fusion entirely.

Why this beats manual batching as a research tool

Manual batching entangles two concerns: the mathematics of one example, and the bookkeeping of the batch axis. vmap factors them, which pays off precisely where bookkeeping gets hard:

  • Per-example gradients: vmap(grad(f)) — the primitive underlying differentially-private SGD (Abadi et al. 2016) and influence-function analysis; in tape-based frameworks this historically required hooks or microbatching.
  • Ensembles: vmap over a parameter axis evaluates $K$ models in one pass (in_axes on the params pytree).
  • Jacobians: jacfwd/jacrev are implemented as vmap of jvp/vjp over basis vectors — the composability is not decorative; core JAX features are built from it.

The idea descends from array-language vectorisation (APL; NESL's flattening transform — Blelloch 1995) and auto-batching literature (Bergstra et al.'s Theano scan being the loop-shaped ancestor). The per-primitive-rule formulation is described in Frostig, Johnson, Leary (2018), Compiling machine learning programs via high-level tracing.

Limits: primitives without batching rules (rare, mostly exotic linear-algebra or IO) fall back or fail loudly; data-dependent shapes per example (ragged work) do not fit the model — padding or masking remains the standard workaround, as everywhere in fixed-shape compilation.

References

  • Abadi et al. (2016), Deep learning with differential privacy.
  • Blelloch (1995), NESL: a nested data-parallel language — the flattening ancestry.
  • Frostig, Johnson, Leary (2018), Compiling machine learning programs via high-level tracing.

What to learn next