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.
Quick wins for a faster PC:
Repair Windows errors before they cause bigger problemsFix Now →Fix the driver behind crashes, sound loss and screen glitchesFind Drivers →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.
#1 Best Overall
What JAX provides
- Array computing:
jax.numpymirrors 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.
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.
Crashes, No Sound, or Screen Glitches?
Random freezes, missing sound and display glitches usually trace back to one bad driver. Find and replace yours safely.Free scan · under a minuteWindows Errors? Fix Them Before They Spread
Repair common Windows errors and clear accumulated junk for a smoother, more stable PC - no reinstall needed.Free scan · no reinstalljax.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.
Rank #2
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.
Do these 3 things before closing this tab:
1Clear out junk files and repair common Windows errors2Fix the driver behind crashes, sound loss and screen glitches3Repair Windows errors before they cause bigger problemsComposing 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.
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.
Recommended Free Tools
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.
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.
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.
Rank #4
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.
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.
Fix: inspect the available devices, reshape the leading axis to match the replica count, or use vmap when you only need within-device batching.
Best Value
- Used Book in Good Condition
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.
The Tool Desk
Outbyte PC Repair FREEClear out junk files and repair common Windows errorsFree Scan →Outbyte Driver Updater FREEFix the driver behind crashes, sound loss and screen glitchesFind Drivers →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.
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.
Quick Recap
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.

