jax-cuda13-pjrt
JAX XLA PJRT Plugin for NVIDIA GPUs
Decision gist · record as of 2026-08-14
Yes, if you have NVIDIA GPU hardware with CUDA 13 and want to use JAX for accelerated computing. The package is actively maintained, permissively licensed, and has no known vulnerabilities. Install friction is moderate due to platform specificity and CUDA system dependency, but necessary for GPU support. Not applicable on CPU-only or non-NVIDIA systems.AI-flagged interpretation of the facts on this page — verify before relying
Before you install
- Requires CUDA 13 runtime libraries installed on the system and an NVIDIA GPU; only available on Linux x86_64 or aarch64 platforms.
- Medium install friction due to platform-specific wheels (manylinux_2_27 x86_64 and aarch64 only) and CUDA 13 system dependency.
- Active maintenance with recent release (29 days old) and strong upstream project health (36159 stars).
License · maintenance · safety
Apache-2.0 (permissive) — Apache-2.0 permissive license allows commercial and private use with minimal restrictions; suitable for most production and research contexts.
last release 2026-07-16 (29 days) · last repo commit 2026-08-14 · 36,159 stars
0 known vulnerabilities (OSV.dev, 2026-08-14) · 415,248 downloads/mo, #6,827 on PyPI
Alternatives
Verify before relying
pip install jax-cuda13-pjrt
import jax
import jax.numpy as jnp
# JAX operations now dispatch to NVIDIA GPU via PJRT
x = jnp.ones((100, 100))
y = jnp.dot(x, x) # runs on GPU- Whether this package is required separately or automatically installed as a JAX dependency on compatible systems
- Specific NVIDIA GPU models and compute capabilities supported by CUDA 13 PJRT plugin
- Performance characteristics or overhead compared to other JAX GPU backends
What it is and what it does
This package is a PJRT (Portable Runtime for JAX) plugin that bridges JAX's numerical computing framework to NVIDIA GPUs via CUDA 13. It enables JAX's core transformations—automatic differentiation (grad), JIT compilation, and vectorization (vmap)—to execute on NVIDIA hardware by interfacing with XLA's compiler infrastructure.
When installed, it allows JAX programs to transparently offload array computations and neural network training to NVIDIA GPUs without code changes. The plugin handles device memory management, kernel dispatch, and multi-device coordination. It is a platform-specific binary distribution (wheels for Linux x86_64 and aarch64 only) that must match your system's CUDA 13 installation.
Use it for
- Train neural networks on NVIDIA GPUs using JAX's automatic differentiation and JIT compilation.
- Accelerate scientific computing workloads (linear algebra, FFTs, stochastic simulations) on GPU hardware.
- Run JAX-based machine learning research code on NVIDIA clusters or cloud instances with CUDA 13.
- Develop high-performance numerical applications requiring GPU parallelism using vmap.
Worth the install?
AI-flagged interpretation of the facts on this page. Verify before relying on it.
Yes, if you have NVIDIA GPU hardware with CUDA 13 and want to use JAX for accelerated computing.
The package is actively maintained, permissively licensed, and has no known vulnerabilities. Install friction is moderate due to platform specificity and CUDA system dependency, but necessary for GPU support. Not applicable on CPU-only or non-NVIDIA systems.
Install
jax-cuda13-pjrt on PyPI
Before you install
Medium install friction due to platform-specific wheels (manylinux_2_27 x86_64 and aarch64 only) and CUDA 13 system dependency. Active maintenance with recent release (29 days old) and strong upstream project health (36159 stars).
Requires CUDA 13 runtime libraries installed on the system and an NVIDIA GPU; only available on Linux x86_64 or aarch64 platforms.
License in practice
Apache-2.0 permissive license allows commercial and private use with minimal restrictions; suitable for most production and research contexts.
Quickstart
pip install jax-cuda13-pjrt
import jax
import jax.numpy as jnp
# JAX operations now dispatch to NVIDIA GPU via PJRT
x = jnp.ones((100, 100))
y = jnp.dot(x, x) # runs on GPU
Verify before relying
- Whether this package is required separately or automatically installed as a JAX dependency on compatible systems
- Specific NVIDIA GPU models and compute capabilities supported by CUDA 13 PJRT plugin
- Performance characteristics or overhead compared to other JAX GPU backends
Package facts
| License | Apache-2.0 permissive |
| Python support | Not specified |
| Install friction | Medium. Platform-specific wheel |
| Runtime dependencies | None |
| Maintenance | Actively maintained 29 days since the last release |
| Last repo commit | |
| First released | |
| Downloads | 415,248 / month, #6,827 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 :: 3Programming Language :: Python :: Free Threading :: 3 - Stable |
Evidence: jax_cuda13_pjrt-0.11.0-py3-none-manylinux_2_27_aarch64.whl; jax_cuda13_pjrt-0.11.0-py3-none-manylinux_2_27_x86_64.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 › “JAX NVIDIA GPU acceleration”
- jax-cuda13-pjrtProvides NVIDIA GPU acceleration for JAX numerical computing via the…
- jax-cuda12-pjrtProvides NVIDIA GPU acceleration for JAX numerical computing via the…
- jax-cuda12-pluginEnables JAX to run numerical computations and machine learning…
Give your agent the search over MCP, or paste the wish link into any chat.
More Scientific/Engineering packages
NumPy provides an N-dimensional array object and a comprehensive suite of mathematical, linear algebra, Fourier transform, and random number functions for scientific computing in Python.
pandas provides fast, flexible data structures (Series and DataFrame) for loading, cleaning, transforming, and analyzing labeled or relational data in Python.
scipy provides numerical algorithms for mathematics, science, and engineering—including optimization, integration, linear algebra, Fourier transforms, signal and image processing, and ODE solvers—built on numpy arrays.
scikit-learn provides a comprehensive Python library for supervised and unsupervised machine learning, including classification, regression, clustering, dimensionality reduction, and model evaluation tools built on NumPy and SciPy.
Install it if you need to train, evaluate, or deploy supervised or unsupervised learning models.
dill extends Python's pickle module to serialize and deserialize a much wider range of Python objects, including functions, lambdas, classes, and interpreter sessions, to byte streams for storage or network transmission.
Multiprocess is an enhanced fork of Python's standard multiprocessing library that uses dill for better serialization, allowing you to spawn processes with a threading-like API and share complex objects between them.
Install it if you use multiprocessing and encounter pickle serialization limits with lambdas or complex objects.
See also jax-cuda12-pjrt · torchax · augmax · tokamax · jax-cuda13-plugin · jax · jax-cuda12-plugin · jaxlib · jmp · klujax