Tooling & Hardware

JAX

JAX is a Google-associated deep learning framework built around function transformations (grad, jit, vmap), popular in research and large-scale training.

JAX (Google-associated) is a deep learning framework built around composable function transformations — grad (automatic differentiation), jit (compilation), vmap (vectorization) — rather than an object-oriented model API like PyTorch. It's popular in research and for some of the largest-scale training runs, since its functional style composes unusually well with distributed hardware, though it has a steeper learning curve than PyTorch.

How it works

JAX asks you to write pure functions over arrays, then transforms them. jit traces the function once with abstract shape/dtype placeholders, lowers the resulting graph to XLA, and compiles a fused kernel for the target device. grad applies reverse-mode automatic differentiation — the same backpropagation math, expressed as a program transformation. vmap adds a batch dimension to code written for a single example, and pmap or the shard_map/jit sharding APIs spread the same computation across devices for distributed training.

Because transformations compose, jit(grad(vmap(f))) is a normal thing to write. The cost of that purity is explicit state: parameters, optimizer state, and random keys are all passed in and returned rather than held as attributes.

When it breaks

  • Recompilation thrash. Compilation is keyed on input shapes, so variable-length batches trigger a fresh XLA compile every step. Padding to fixed buckets is the standard workaround.
  • Tracers leak into Python. Inside jit, values are abstract, so if x > 0 or print(x) on a traced array fails or does nothing useful. Control flow must use lax.cond and lax.scan.
  • Side effects vanish. Mutating a global or appending to a list happens only during the trace, not on subsequent calls — a particularly quiet source of wrong metrics.
  • Randomness is explicit. Reusing a PRNG key instead of splitting it gives identical "random" draws, which can silently break dropout or initialization.

See also: PyTorch, GPU

Learn more: Tooling & The Dev Stack · JAX official docs

Mentioned in

Lessons where this comes up in context.

On this page