On this page
Overview
JAX is a Python library for high-performance numerical computing and large-scale machine learning. Functions written with a NumPy-style array API can be transformed with grad for automatic differentiation, jit for XLA compilation, and vmap for automatic vectorization. Its execution model also extends from CPUs to supported GPUs, TPUs, and multi-device systems.
Features and best fit
Based on official documentation; not hands-on tested · Content checked:
Key features
Compose differentiation, compilation, and vectorization on the same function
jax.grad differentiates functions written in Python and a NumPy style, including functions with loops, branches, recursion, and closures. jax.jit compiles functions through XLA, while jax.vmap maps a function over array axes without requiring manual batch-dimension plumbing. These transformations compose, allowing patterns such as compiled, vectorized per-example gradients.
Sources: [1]
Best fit
Fits workloads that combine gradients with accelerator-oriented numerical computing
JAX fits machine learning, optimization, and scientific-computing workloads that need automatic differentiation while moving the same computation across CPUs, GPUs, or TPUs. Automatic partitioning, explicit sharding, and manual per-device programming provide paths from single-device experiments to research and engineering across many devices.
Sources: [1]
Before adoption
Validate JIT semantics and hardware-specific installation before adoption
The project describes itself as a research project rather than an official Google product and explicitly warns of sharp edges. Normal Python control flow is constrained under jax.jit, so existing NumPy code is not automatically a drop-in optimized workload. Support and installation packages differ across CPUs, NVIDIA GPUs, TPUs, AMD GPUs, Intel GPUs, and operating systems, with some combinations marked experimental. Stable release 0.11.2 includes breaking changes, so upgrades require changelog review. The project is Apache-2.0 licensed.
Official sources
- [1]jax-ml/jax — README(2026-10-04)
- [2]JAX 0.11.2 release(2026-10-04)
- [3]jax-ml/jax — LICENSE(2026-10-04)
Supplemental curator note
The key comparison point is not any single feature but the ability to compose transformations such as grad, jit, and vmap. Before adopting it, test the target OS and accelerator, Python control flow under JIT, and the numerical behavior your workload must reproduce.
Try it in 3 steps
- 1
Create and activate a virtual environment
Use Python 3.12 or newer as required by JAX 0.11.2. This POSIX-shell command isolates the check from the system environment.
python3 -m venv .venv && . .venv/bin/activate - 2
Install the stable CPU build
For a GPU or TPU, choose the hardware-specific extra from the official installation table instead.
python -m pip install --upgrade 'jax==0.11.2' - 3
Check automatic differentiation and array operations
A derivative value of 6.0 and a JAX array confirm the minimal CPU path.
python -c "import jax; import jax.numpy as jnp; print(jax.grad(lambda x: x**2)(3.0)); print(jnp.arange(3))"
Growth
Growth trends · Last 30 days
36,374 Stars
Trend data is still being collected.
Development activity
Last 90 days · weekly
- Commits (last 30 days)
- 671
- Open PRs
- 862
Development activity is still being collected.
Built with
Categories and tags
GitHub data
GitHub dataView detailed GitHub data
GitHub Topics
- jax
- Stars
- 36,374
- Forks
- 3,834
- Watchers
- 333
- Open issues
- 1,778
- Contributors
- 396
- Owner type
- Organization
- Primary language
- Python
- License
- Apache-2.0
- Repository last updated
- Oct 4, 2026
Related information
Write a related articleShare a guide or use case for this OSS in Markdown. Articles are published after administrator approval.
Explore next
- CatBoost9,134 Stars
3 shared tag(s) · 3 shared category(s)
train classification, regression, and ranking models with gradient-boosted decision trees that handle numerical and categorical features
C++ - scikit-learn67,463 Stars
3 shared tag(s) · 2 shared category(s) · Same language
a Python machine-learning library for composing preprocessing, training, model selection, and evaluation through estimators and pipelines
Python - SciPy15,075 Stars
3 shared tag(s) · 2 shared category(s) · Same language
Add optimization, integration, linear algebra, statistics, FFTs, and signal processing to NumPy arrays with a scientific-computing library
Python - Keras 364,347 Stars
3 shared tag(s) · 1 shared category(s) · Same language
use JAX, TensorFlow, and PyTorch through one high-level model API
Python - Streamlit45,893 Stars
3 shared tag(s) · 1 shared category(s) · Same language
Turn Python scripts into interactive data apps with widgets, dataframes, charts, multipage navigation, and chat
Python - bitsandbytes8,510 Stars
3 shared tag(s) · 1 shared category(s) · Same language
quantize PyTorch models to 8-bit and 4-bit formats to reduce memory use for LLM inference and QLoRA training
Python
Report incorrect information
Tell us if any listing information is incorrect or outdated.