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 learningDeep LearningDeep learning is machine learning using multi-layer neural networks, which learn their own features from raw data instead of relying on hand-engineered ones.
framework built around composable function transformations — grad
(automatic differentiation), jit (compilation), vmap
(vectorization) — rather than an object-oriented model API like
PyTorchPyTorchPyTorch is the dominant deep learning framework in both research and production, providing tensor computation, GPU dispatch, and automatic differentiation.. 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 backpropagationBackpropagationBackpropagation is the algorithm that computes the gradient of a neural network's loss with respect to every parameter, by applying the chain rule backward through the network.
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 trainingDistributed Training (FSDP, DeepSpeed, Megatron-LM)Distributed training frameworks split a model's parameters, gradients, and optimizer state across many GPUs so models too large for one GPU can still be trained..
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, soif x > 0orprint(x)on a traced array fails or does nothing useful. Control flow must uselax.condandlax.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: PyTorchPyTorchPyTorch is the dominant deep learning framework in both research and production, providing tensor computation, GPU dispatch, and automatic differentiation., GPUGPU (Graphics Processing Unit)GPUs, originally built for rendering graphics, turned out to be extremely well-suited to the parallel matrix multiplications deep learning requires.
Learn more: Tooling & The Dev Stack · JAX official docs
Mentioned in
Lessons where this comes up in context.
TensorFlow
TensorFlow is a deep learning framework, dominant in the mid-2010s, still widely used in production and edge deployment (via TensorFlow Lite).
Hugging Face
Hugging Face is an ecosystem — transformers, datasets, and the Hub — providing pretrained models and standardized tooling on top of frameworks like PyTorch.