OSS TanbouSign in with GitHub

Compose automatic differentiation, JIT compilation, vectorization, and distributed execution for Python and NumPy-style programs

About these scores

OSS scale score is an unbounded metric that log-compresses and weights Stars, Watchers, Forks, and Contributors. Discovery score is the current OSS scale score minus the score at discovery. Update pace is commits in the last 30 days, growth momentum is the OSS scale score difference within the recent observation window, and OSS health is a 0–100 rating based on available recency, Community Health, and release data.

Stars
36,374
Primary language
Python
License
Apache-2.0
Repository last updated
Oct 4, 2026
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.

Sources: [1][2][3]

Official sources

  1. [1]jax-ml/jax — README(2026-10-04)
  2. [2]JAX 0.11.2 release(2026-10-04)
  3. [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. 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. 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. 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))"
Check the official README

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
Write a related article

Share a guide or use case for this OSS in Markdown. Articles are published after administrator approval.

Report incorrect information

Tell us if any listing information is incorrect or outdated.

After reading this page, do you know what to do next?