Math & training · Glossary term

What is JAX?

A Python library for transforming numerical functions with automatic differentiation, compilation, vectorization, and parallel execution across accelerators. Its transformations work best with explicit state and functional-style code.

What people say

“A NumPy-like system for accelerated machine learning.”

What is the common confusion about JAX?

JAX does not prohibit all stateful programming, but hidden mutation inside transformed functions can produce incorrect or unsupported behavior.

Learn JAX in the course

Start with

  • Introduction to JAX

    PyTorch mutates tensors. TensorFlow builds graphs. JAX compiles pure functions. That last one changes how you think about deep learning.

    Phase 03: Deep Learning Core

Lessons that name JAX in a title or section

  • Introduction to PyTorch

    You built the engine from pistons and crankshafts. Now learn the one everyone actually drives. Build and train neural networks using PyTorch's nn.Module, nn.Sequential, and autograd.

    Phase 03: Deep Learning Core

Taught in Phase 03: Deep Learning Core.

  • AutogradA system that records or transforms tensor operations so it can compute derivatives, usually with reverse-mode automatic differentiation.
  • TensorA typed array with a shape, data type, and device placement that frameworks use to represent inputs, parameters, activations, and gradients.
  • CUDANVIDIA's platform and programming model for general-purpose computation on compatible GPUs.

Sources

More terms in Math & training

Open the Math & training list in the glossary

This entry comes from glossary/terms.md on GitHub. Browse all 250 glossary terms.