jax-ml/jax

Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more

View on GitHub ↗Jump to charts ↓Open shareable report →

Data as of . Signed-in members get hourly updates — create a free account.

Summary Information

Updated 2 hours ago
Added to GitGenius on October 13th, 2024
Created on October 25th, 2018
Open Issues & Pull Requests: 2,619 (+1)
GitHub issues: Enabled
Number of forks: 3,842
Total Stargazers: 36,390 (+1)
Total Subscribers: 333 (+0)

Repository Insights (GitGenius)

Median issue/PR response: 10.1 hours
Mean response time: 179.7 days
90th percentile: 782.2 days
Tracked items: 2,731

Maintainer activity

38 people did triage or write work on this repository in the last 12 months.

At least 39% of jax's 38 maintainers work at Google. 25 say where they work, and 15 of those are Google.

Counts unlabeled, assigned, unassigned, milestoned, demilestoned, locked, unlocked over the last 12 months. These are issue and pull request events that require triage or write permission. Commits and code review are not counted. labeled and renamed are excluded because GitHub issue forms record the issue author as the actor. Figures from October 7, 2026. This count is not comparable across projects: each project's automation decides which of these events a person emits.

How this project is maintained

About 14% of issues opened in the past year have never received a reply. 87% of open issues come from outside the core team, so the backlog reflects real-world use rather than internal planning. Work labelled "AMD GPU" is answered fastest, typically in under an hour, while "P3 (no schedule)" waits about 13 months. 55% of tracked open issues have had no activity in three months. Only 48% of issues opened in the past year have been closed.

Charts & Analytics

Fetching additional details & charts...

Issue Activity (beta)

Open issues: 1,618
New in 7 days: 11
Closed in 7 days: 33
Avg open age: 776 days
Stale 30+ days: 1,543
Stale 90+ days: 1,369

Recent activity

Opened in 7 days: 11
Closed in 7 days: 29
Comments in 7 days: 9
Events in 7 days: 27

Top labels

  • bug (3,517)
  • enhancement (1,328)
  • question (348)
  • NVIDIA GPU (239)
  • documentation (198)
  • pallas (142)
  • needs info (134)
  • P2 (eventual) (104)

Most active issues this week

Sign in to see which issues are moving.
Sign in

Detailed Description

JAX is a Python library for accelerator-oriented array computation and program transformation, designed for high-performance numerical computing and large-scale machine learning. The library enables composable transformations of Python and NumPy programs, allowing developers to differentiate, vectorize, and compile code to run on GPUs, TPUs, and other hardware accelerators. JAX uses XLA as its compilation backend to scale numerical programs across thousands of devices.

The core of JAX is an extensible system for transforming numerical functions. The jax.grad function provides automatic differentiation capabilities, supporting both reverse-mode differentiation (backpropagation) and forward-mode differentiation that can be composed arbitrarily to any order. Differentiation works through loops, branches, recursion, and closures, enabling derivatives of derivatives of derivatives. The jax.jit function compiles pure functions end-to-end using XLA, while jax.vmap provides auto-vectorization by mapping functions along array axes and pushing loops down onto primitive operations for better performance. These transformations can be composed together, allowing developers to obtain efficient Jacobian matrices or per-example gradients by combining vmap with grad and jit.

For scaling computations across multiple devices, JAX offers three approaches: compiler-based automatic parallelization where the compiler determines data sharding and computation partitioning, explicit sharding with automatic partitioning where data shardings are visible in JAX types, and manual per-device programming with explicit collectives for fine-grained control. The library supports multiple platforms including Linux x86_64, Linux aarch64, Mac aarch64, and Windows, with varying levels of support for CPU, NVIDIA GPU, Google TPU, AMD GPU, Apple GPU, and Intel GPU backends.

The repository is classified across multiple domains including machine learning, differentiable programming, numerical computing, automatic differentiation, and high-performance computing, reflecting its broad applicability across research and production machine learning workloads.

JAX is explicitly positioned as a research project rather than an official Google product, with documentation acknowledging sharp edges and encouraging community feedback through bug reports and feature discussions. The library provides comprehensive installation instructions for different hardware configurations and maintains detailed reference documentation covering both user-facing APIs and developer guidelines for contributing to the project.