jax
Differentiate, compile, and transform Numpy code.
Decision gist · record as of 2026-08-14
Yes. JAX is production-stable (Development Status 5), actively maintained, has no known vulnerabilities, and installs with low friction. It is the right choice if you need automatic differentiation, XLA compilation, or large-scale distributed training. Install it if your workflow involves numerical computing, machine learning research, or scientific computing on accelerators; avoid it only if you need only basic NumPy operations without transformation capabilities.AI-flagged interpretation of the facts on this page — verify before relying
Before you install
- Requires Python 3.12 or later.
- GPU/TPU acceleration requires jaxlib and appropriate hardware drivers; CPU-only use works out of the box.
- Low friction installation with a pure Python wheel.
License · maintenance · safety
Apache-2.0 (permissive) — Licensed under Apache-2.0 (permissive), allowing use in commercial and proprietary projects with minimal restrictions—attribution required but no copyleft obligations.
last release 2026-07-16 (29 days) · last repo commit 2026-08-14 · 36,157 stars
0 known vulnerabilities (OSV.dev, 2026-08-14) · 23,495,676 downloads/mo, #942 on PyPI
Alternatives
Verify before relying
import jax
import jax.numpy as jnp
def f(x):
return jnp.sum(x ** 2)
grad_f = jax.grad(f)
result = grad_f(jnp.array([1.0, 2.0, 3.0]))- Whether the package's 'sharp edges' documented in its gotchas notebook materially affect typical machine learning workflows.
- Performance characteristics and memory overhead compared to alternatives for specific hardware targets (GPU vs TPU vs CPU).
- Stability of experimental features (Apple GPU, Intel GPU, AMD GPU on some platforms) listed in the installation matrix.
What it is and what it does
JAX is a numerical computing library that transforms Python and NumPy functions into differentiable, compilable, and parallelizable code. It provides three core transformations: automatic differentiation (grad), just-in-time compilation (jit), and auto-vectorization (vmap), which can be composed arbitrarily. The library uses XLA to compile and scale computations across accelerators like GPUs and TPUs.
Typically used in machine learning research and production systems, JAX lets you write numerical code once and apply transformations to get gradients, compiled kernels, or batched operations without rewriting. It supports reverse-mode differentiation (backpropagation), forward-mode differentiation, and higher-order derivatives. Scaling modes range from automatic compiler-based parallelization to explicit per-device programming for distributed training.
Use it for
- Training neural networks with automatic gradient computation and JIT compilation for speed.
- Computing Jacobians and per-example gradients efficiently via vmap composition with grad.
- Scaling machine learning workloads across multiple GPUs or TPUs using explicit or automatic sharding.
- Prototyping differentiable algorithms where you need derivatives of complex control flow (loops, branches, recursion).
- Scientific computing with automatic differentiation for physics simulations and optimization.
Worth the install?
AI-flagged interpretation of the facts on this page. Verify before relying on it.
Yes.
JAX is production-stable (Development Status 5), actively maintained, has no known vulnerabilities, and installs with low friction. It is the right choice if you need automatic differentiation, XLA compilation, or large-scale distributed training. Install it if your workflow involves numerical computing, machine learning research, or scientific computing on accelerators; avoid it only if you need only basic NumPy operations without transformation capabilities.
Install
jax on PyPI
Before you install
Low friction installation with a pure Python wheel. Active maintenance with recent releases; the repository shows strong community engagement (36157 stars) and current Python 3.12+ support. Runtime dependencies are standard numerical libraries (numpy, scipy, opt_einsum, ml_dtypes, jaxlib).
Requires Python 3.12 or later. GPU/TPU acceleration requires jaxlib and appropriate hardware drivers; CPU-only use works out of the box.
License in practice
Licensed under Apache-2.0 (permissive), allowing use in commercial and proprietary projects with minimal restrictions—attribution required but no copyleft obligations.
Quickstart
import jax
import jax.numpy as jnp
def f(x):
return jnp.sum(x ** 2)
grad_f = jax.grad(f)
result = grad_f(jnp.array([1.0, 2.0, 3.0]))
Verify before relying
- Whether the package's 'sharp edges' documented in its gotchas notebook materially affect typical machine learning workflows.
- Performance characteristics and memory overhead compared to alternatives for specific hardware targets (GPU vs TPU vs CPU).
- Stability of experimental features (Apple GPU, Intel GPU, AMD GPU on some platforms) listed in the installation matrix.
Package facts
| License | Apache-2.0 permissive |
| Python support | Supports the current Python release >=3.12 |
| Install friction | Low. Pure-Python wheel |
| Runtime dependencies | 5 packagesjaxlibml_dtypesnumpyopt_einsumscipy |
| Maintenance | Actively maintained 29 days since the last release |
| Last repo commit | |
| First released | |
| Downloads | 23,495,676 / month, #942 on PyPI 30-day window, as of 2026-08-14 |
| Known vulnerabilities | None known OSV.dev, checked 2026-08-14 |
| Classifiers | Development Status :: 5 - Production/StableProgramming Language :: Python :: 3.12Programming Language :: Python :: 3.13Programming Language :: Python :: 3.14Programming Language :: Python :: Free Threading :: 3 - Stable |
Evidence: jax-0.11.0-py3-none-any.whl
Tags
Let your AI agent find packages like this
Example. Real query, live index.
You found this page by searching. An agent finds it by wishing: SkillFed indexes 14,416 PyPI packages by what they can do, searchable in plain language.
wish › “jit compilation numpy”
- jaxJAX is a Python library for automatic differentiation, XLA…
- nvidia-nvjitlink-cu12Provides NVIDIA's JIT LTO compiler library for CUDA 12, enabling…
- numbaggNumbagg provides fast N-dimensional aggregation and moving-window…
Give your agent the search over MCP, or paste the wish link into any chat.
More Artificial Intelligence packages
LiteLLM provides a unified Python interface to call 100+ LLM providers (OpenAI, Anthropic, Gemini, Bedrock, Azure, and others) using OpenAI-compatible API format, available as both a Python SDK and a self-hosted AI Gateway proxy server.
Install it if you need to work with multiple LLM providers or want to centralize LLM routing in your organization.
Client library and CLI tool for downloading, uploading, and managing models, datasets, and repositories on the Hugging Face Hub platform.
Install it if you work with Hugging Face Hub models or datasets.
LangChain provides a framework for building agents and LLM-powered applications by composing language models, tools, and memory through a unified API that abstracts over multiple model providers.
hf-xet provides chunk-based deduplication and efficient file transfer for the Hugging Face Hub, enabling faster uploads and downloads of large files with local disk caching.
Tokenizers converts raw text into token sequences for NLP models, with support for training custom vocabularies and using pre-built tokenizers (BPE, WordPiece) optimized for speed via Rust.
Transformers provides a unified framework for loading, fine-tuning, and running state-of-the-art pretrained models across text, vision, audio, video, and multimodal tasks using PyTorch, JAX, or TensorFlow.
Install it if you need to run or train any transformer-based model for NLP, vision, audio, or multimodal tasks.
See also grain · jax-cuda12-plugin · jax-jumpy · jaxellip · jaxlib · augmax · distrax · klujax · chex · jax-cuda13-pjrt