Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.

Some links on this page are affiliate links: if you buy through them we may earn a commission, at no extra cost to you.

JAX is a Python library for accelerator-oriented array computing and program transformation. It gives you a NumPy-style API through jax.numpy (usually imported as jnp), then lets you compile, differentiate, batch and parallelize numerical code for CPUs, NVIDIA GPUs and Google TPUs. Its core compiler path uses XLA, so one mostly functional program can target different accelerators without being rewritten for each device.

This guide explains JAX’s programming model, the roles of jax.jit, jax.grad, jax.vmap and jax.pmap, backend installation, comparisons with NumPy, PyTorch and TensorFlow, and the practical limitations that affect research and production work.

What Google JAX is

JAX is software, not a hardware product or a hosted service. It combines a NumPy-inspired array interface with transformations that rewrite numerical Python functions. The same conceptual computation can run locally on a CPU, on an NVIDIA or AMD accelerator where supported, or on Google Cloud TPU hardware.

Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.

JAX arrays look familiar to NumPy users, but they are immutable and designed to be traced and compiled. You normally express calculations as functions that accept arrays and return arrays. That style allows JAX to inspect the computation, transform it and hand it to the Open XLA compiler.

What JAX provides

  • Array computing: jax.numpy mirrors much of the NumPy API.
  • Automatic differentiation: gradients can be generated from numerical functions.
  • Compilation: eligible functions can be compiled into optimized backend code.
  • Vectorization: a function written for one example can be lifted over a batch.
  • Parallel execution: replicated computations can run across multiple XLA devices.

JAX is therefore a foundation for machine-learning research, optimization, simulation and scientific programs that benefit from differentiability, vectorized work, compilation or accelerator scaling. Neural-network, optimizer, probabilistic and deployment libraries are commonly built on top of this core rather than being part of the minimal array API itself.

How JAX executes a function

When a transformed function runs, JAX traces the operations performed with JAX values and builds an intermediate representation. jax.jit sends that representation to XLA, which can fuse operations and generate code for the selected backend. Compilation is cached using input types and related compilation conditions. Consequently, the first call can be slower than later calls because it includes compilation; later calls can reuse the compiled executable when their requirements match.

Why the programming model matters

Tracing works best when a function is mostly pure: its result should be determined by its arguments, and it should avoid hidden mutation or side effects that the compiler cannot represent. Python control flow that depends on ordinary runtime values may need to be expressed with JAX-compatible control-flow operations or arranged so the relevant values are known during tracing. Changing shapes or dtypes can require another compilation, so stable batch and feature dimensions are important for throughput.

Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.

The four transformations you will use most

jax.jit: just-in-time compilation

jax.jit takes a Python function, traces its JAX operations and compiles the resulting computation. Use it around a numerically substantial function that will be called repeatedly.

import jax
import jax.numpy as jnp

def energy(x):
    return jnp.sum(jnp.sin(x) ** 2)

compiled_energy = jax.jit(energy)
value = compiled_energy(jnp.ones((1024,)))

The first invocation may include tracing and compilation. Reusing the compiled function with compatible argument structures lets subsequent calls avoid that setup. JIT compilation is not automatically faster for every tiny operation; compilation overhead and data-transfer costs can outweigh gains for short, irregular work.

jax.grad: automatic differentiation

jax.grad transforms a scalar-output numerical function into a function that computes its gradient with respect to a chosen argument. Differentiation can be composed with other transformations.

import jax
import jax.numpy as jnp

def objective(w, x, target):
    prediction = jnp.dot(x, w)
    return jnp.mean((prediction - target) ** 2)

grad_objective = jax.grad(objective)
g = grad_objective(
    jnp.zeros((3,)),
    jnp.ones((8, 3)),
    jnp.ones((8,)),
)

JAX supports numerical differentiation and composition of differentiation with batching and compilation. The function being differentiated must use operations for which JAX has differentiation rules.

Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.

jax.vmap: automatic vectorization

jax.vmap lifts a function written for one example so it operates efficiently over a batch. Instead of manually adding batch dimensions to every operation, you describe the single-example function and map it.

import jax
import jax.numpy as jnp

def score_one(w, x):
    return jnp.dot(x, w)

score_batch = jax.vmap(score_one, in_axes=(None, 0))
 scores = score_batch(
    jnp.ones((4,)),
    jnp.ones((16, 4)),
)

The mapped axis is specified with in_axes; None keeps an argument shared across the batch. This is different from writing a separate loop in Python: the mapped computation remains available to JAX’s transformations.

jax.pmap: replicated multi-device execution

jax.pmap compiles a function and executes replicas in parallel on multiple XLA devices, such as GPU devices or TPU cores. It is intended for multi-device parallel execution. The leading mapped axis normally corresponds to the replicas, and the input must provide enough slices for the available devices.

import jax
import jax.numpy as jnp

def per_device_sum(x):
    return jnp.sum(x)

parallel_sum = jax.pmap(per_device_sum)
# The leading dimension supplies one slice per participating device.
# result = parallel_sum(batch_shaped_for_device_count)

pmap and vmap solve different problems. vmap vectorizes within array operations, usually for examples in one process and device. pmap replicates a computation across physical XLA devices. For newer sharding designs, JAX also exposes broader automatic-parallelization and sharding concepts, but pmap remains the direct primitive for replicated execution.

Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.

Composing transformations

Transformations are designed to compose. A common pattern is to differentiate one-example loss, vectorize that gradient over a batch and compile the result:

import jax
import jax.numpy as jnp

def loss_one(w, x, y):
    prediction = jnp.dot(x, w)
    return (prediction - y) ** 2

grad_one = jax.grad(loss_one)
grad_batch = jax.vmap(grad_one, in_axes=(None, 0, 0))
compiled_grad_batch = jax.jit(grad_batch)

weights = jnp.zeros((3,))
features = jnp.ones((32, 3))
targets = jnp.ones((32,))
gradients = compiled_grad_batch(weights, features, targets)

This works when the functions remain traceable and mostly free of side effects.

JAX compared with NumPy, PyTorch and TensorFlow

No single framework is best for every workload. The practical differences are the programming model, compilation behavior, differentiation, scaling, ecosystem and installation requirements.

Axis JAX NumPy PyTorch TensorFlow
Programming model NumPy-style arrays plus composable functional transformations Direct numerical array operations Imperative/object-oriented training APIs are common Imperative and graph-oriented APIs are available
Compilation Optional JIT tracing through XLA; cache reuse depends on input types and compilation conditions No JAX-style transformation layer Compilation facilities exist, but the workflow and tracing rules differ Graph and compilation workflows differ from JAX’s transformation model
Differentiation Automatic differentiation that composes with batching and compilation Not a built-in automatic-differentiation framework Automatic differentiation is integrated with tensor and training workflows Automatic differentiation is integrated with tensor and training workflows
Scaling target CPU, supported GPU and TPU backends; multi-device execution and sharding concepts Primarily host numerical computing unless paired with other tools GPU-centered machine-learning workflows with its own distributed APIs GPU/TPU machine-learning workflows with its own distributed APIs
Ecosystem emphasis Core numerical foundation used by higher-level ML and scientific libraries Broad general-purpose Python numerical ecosystem Large end-to-end deep-learning ecosystem Large end-to-end deep-learning ecosystem
Setup Backend-specific wheels, drivers and plugins matter Usually simpler CPU installation Accelerator installation depends on the selected build and drivers Accelerator installation depends on the selected build and drivers

Choose JAX when your program benefits from a NumPy-like surface combined with transformations, differentiability and XLA execution. Choose plain NumPy when you need straightforward CPU array work without tracing. PyTorch or TensorFlow may be a better fit when your team depends on their established model, data-loading, deployment or training ecosystems. These are workflow trade-offs, not universal speed rankings: performance depends on backend, array shapes, compilation, memory movement and the workload itself.

What’s actually slowing this PC down?

Pick the symptom - the matching free tool is one click away.

Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.

Hardware support and installation

The jax package supplies the Python API, while jaxlib contains compiled binaries and backend support. Install the variant that matches the hardware and operating system you actually use.

Target Documented command Important qualification
CPU pip install -U jax Supported Linux, macOS and Windows combinations include Linux x86_64, Linux aarch64, Apple ARM macOS and Windows x86_64, with platform caveats.
NVIDIA GPU pip install -U "jax[cuda13]" CUDA 13 wheels are documented for Linux; Windows WSL2 support is experimental.
AMD GPU pip install -U "jax[rocm7-local]" ROCm must already be installed. Linux is the primary target; WSL2 support is experimental.
Google Cloud TPU VM pip install "jax[tpu]" Use a Linux TPU VM environment; installing the package does not create or allocate a TPU.

Apple GPU acceleration is not supported by JAX’s documented installation path, so Apple users should use the CPU installation unless they are working in a separately supported environment. Intel GPU support is experimental. Backend availability and wheel names can change, so verify the current platform matrix before pinning a production environment.

Verify the backend from Python

import jax
print(jax.devices())

The returned device list tells you which backend the current installation can see. If it shows only a CPU when you expected an accelerator, check the wheel, driver or plugin installation and the environment in which Python is running.

Writing JAX-friendly code

Keep transformed functions explicit

Pass parameters and arrays as arguments instead of reading mutable global state. Return updated values rather than mutating arrays in place; JAX arrays are immutable by design.

Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.

Watch shapes, dtypes and static values

Compilation is specialized to input types and related conditions. Feeding many different shapes or dtypes can create multiple compiled executables and increase startup cost. Keep batch dimensions consistent where possible, and separate configuration values that must be known during tracing from data that changes every call.

Account for host-device transfers

Moving data between Python and an accelerator can erase the benefit of a fast kernel. Keep work on the selected device, avoid converting arrays to ordinary NumPy values inside a hot loop, and measure complete steps rather than a single arithmetic operation.

Plan random state and program state

JAX transformations work best when state is represented explicitly in function arguments and return values. Treat random keys and model state as data passed through the computation instead of relying on hidden mutation.

Performance, reliability and cost considerations

  • Warm-up: include the first-call compilation cost when measuring latency, then measure steady-state calls separately.
  • Compilation cache behavior: compatible types and conditions can reuse a compiled executable; shape or dtype changes may trigger recompilation.
  • Workload fit: large, repeated numerical kernels generally give compilation more opportunity to help than tiny one-off calls.
  • Backend variance: CPU, GPU and TPU results differ with hardware, memory limits, operation support and array shapes. The project documentation does not establish a universal speedup percentage.
  • Operational cost: JAX itself is a software library. TPU or GPU usage can introduce separate infrastructure charges through the service or hardware provider you choose.
Independent reader supportYour contribution helps us test, update, and keep practical guides available for everyone.Support on Ko-Fi

Troubleshooting common JAX problems

Only a CPU device appears

Cause: the CPU wheel is installed, the accelerator driver is missing, or the plugin and runtime versions do not match.

Free tools Windows power users keep installed

One-click scans. No signup required.

Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.

Fix: reinstall the documented accelerator extra in the intended virtual environment, confirm the NVIDIA CUDA or AMD ROCm prerequisites, then inspect jax.devices() again. On TPU, run inside the supported Linux TPU VM environment.

The first call is unexpectedly slow

Cause: tracing and XLA compilation happen before the first result.

Fix: perform a warm-up call before latency measurement, reuse the compiled function, and avoid changing shapes or dtypes between calls unless that specialization is intentional.

JIT fails around Python logic or side effects

Cause: the function depends on runtime Python values, hidden mutable state or an operation that cannot be represented in the traced computation.

Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.

Fix: make data dependencies explicit, use JAX-compatible array control flow, move logging and other side effects outside the transformed function, and test the unjitted function first.

Compilation happens repeatedly

Cause: argument structures, shapes, dtypes or static compilation conditions are changing.

Fix: standardize batch shapes and dtypes, keep configuration stable, and reuse one transformed function rather than constructing new wrappers inside a loop.

Multi-device mapping raises a device-count error

Cause: pmap needs one mapped slice per participating XLA device.

Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.

Fix: inspect the available devices, reshape the leading axis to match the replica count, or use vmap when you only need within-device batching.

AMD installation cannot find ROCm

Cause: the JAX ROCm extra expects ROCm to be installed separately.

Fix: install and verify the supported ROCm runtime first, then install jax[rocm7-local] in the same environment and check device discovery.

Is JAX a good choice for research and scientific computing?

JAX is a strong choice when the central algorithm is numerical and benefits from automatic differentiation, vectorized batches, compilation or accelerator execution. Its composable transformations make it possible to express one mathematical function and derive batched, differentiated and compiled forms without maintaining separate implementations.

Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.

It is less attractive when the workload is dominated by irregular Python control flow, opaque side effects, constantly changing shapes or libraries that cannot participate in tracing. In those cases, plain NumPy or an imperative framework may require less restructuring. Teams should prototype the critical kernel, measure warm and steady-state behavior on the intended backend, and confirm that the surrounding ecosystem supports their model, optimizer, data and deployment requirements.

Or skip the browser setup

If you need clean screenshots of rendered JAX notebooks, documentation or dashboards for reports, ScreenshotNeo provides a website screenshot API and MCP server. It removes cookie-consent banners, newsletter popups and chat widgets before capture; bot checks, blank pages, timeouts, failed loads and cache hits are not billed, and response headers identify the page verdict and billing status. AI agents can call its take_screenshot, get_page_info and capture_pdf MCP tools.

One GET request is enough:

curl -G "https://api.screenshotneo.com/v1/shot" -d access_key=YOUR_API_KEY --data-urlencode url=https://screenshotneo.com/docs/ -o shot.webp

See the ScreenshotNeo API documentation for the other 63 options, including full-page lazy-image loading, CSS-selector element capture, dark mode, device and retina settings, PDF output, custom CSS and JavaScript, clicks, waits, blocking rules, headers, cookies, geolocation, transparent backgrounds, resizing, TTL caching, signed links, asynchronous webhooks, bulk capture and usage reporting. The free plan includes 1,000 screenshots per month with no card; paid plans start at $5 for 3,000 shots. Create a free ScreenshotNeo account.

Frequently Asked Questions

Does installing JAX provide access to a GPU or TPU?

No. Installation supplies Python packages and backend support; you still need compatible hardware or a supported cloud environment, plus the required drivers or runtime.

Special offer. See more information about Outbyte and uninstall instructions. Please review EULA and Privacy policy.

Should I use vmap or pmap for a batch?

Use vmap for efficient batching within array operations on a device. Use pmap when you need replicated execution across multiple XLA devices and can provide one mapped slice per device.

Is there a universal JAX speed advantage over other frameworks?

No. Results depend on backend, shapes, compilation overhead, memory movement and workload. Measure the complete program on the hardware and software versions you plan to deploy.

Product prices and availability are accurate as of the date/time indicated and are subject to change. Any price and availability information displayed on Amazon at the time of purchase will apply.