jax-cuda12-plugin
JAX Plugin for NVIDIA GPUs
Decision gist · record as of 2026-08-14
Yes, if you have an NVIDIA GPU with CUDA 12 support and want to use JAX for accelerated numerical computing or machine learning. The plugin is actively maintained, has no known vulnerabilities, and carries a permissive license. Install friction is moderate due to platform specificity and compiled dependencies, but this is typical for GPU acceleration packages. Not applicable on systems without NVIDIA GPUs or on unsupported platforms.AI-flagged interpretation of the facts on this page — verify before relying
Before you install
- Requires NVIDIA GPU with CUDA 12 support and Linux x86_64 or aarch64 platform; Python 3.12 or later.
- Medium install friction due to platform-specific wheels for Linux x86_64 and aarch64 only, plus a compiled runtime dependency on jax-cuda12-pjrt.
- Actively maintained with recent releases and high repository activity (36158 stars).
License · maintenance · safety
Apache-2.0 (permissive) — Apache-2.0 permissive license allows commercial and private use with minimal restrictions.
last release 2026-07-16 (29 days) · last repo commit 2026-08-14 · 36,158 stars
0 known vulnerabilities (OSV.dev, 2026-08-14) · 782,417 downloads/mo, #5,078 on PyPI
Alternatives
Verify before relying
pip install jax-cuda12-plugin
import jax
import jax.numpy as jnp
# GPU acceleration is automatically available for JAX operations
result = jax.jit(lambda x: jnp.dot(x, x))(jnp.ones((100, 100)))- Whether CUDA 12 runtime libraries must be pre-installed on the system for the plugin to function.
- Performance characteristics and memory overhead compared to CPU-only JAX execution.
- Compatibility with specific NVIDIA GPU architectures and CUDA driver versions.
What it is and what it does
jax-cuda12-plugin is a runtime plugin that bridges JAX (a numerical computing library for automatic differentiation and XLA compilation) to NVIDIA GPUs via CUDA 12. When installed, it registers GPU support as a backend for JAX, allowing numerical operations and machine learning computations to execute on NVIDIA hardware instead of CPU.
The plugin depends on jax-cuda12-pjrt, a compiled component that handles low-level GPU communication. It is distributed as platform-specific wheels for Linux x86_64 and aarch64 only, and requires Python 3.12 or later. Installation is straightforward via pip, but GPU availability and proper CUDA environment setup are prerequisites for runtime use.
Use it for
- Accelerate machine learning training by offloading gradient computation and matrix operations to NVIDIA GPUs.
- Speed up numerical simulations and scientific computing workloads using JAX's automatic differentiation and JIT compilation.
- Enable per-example gradient computation and vectorized operations on GPU for batch processing in deep learning pipelines.
- Scale JAX programs across multiple GPUs using data parallelism and sharding strategies on NVIDIA hardware.
Worth the install?
AI-flagged interpretation of the facts on this page. Verify before relying on it.
Yes, if you have an NVIDIA GPU with CUDA 12 support and want to use JAX for accelerated numerical computing or machine learning.
The plugin is actively maintained, has no known vulnerabilities, and carries a permissive license. Install friction is moderate due to platform specificity and compiled dependencies, but this is typical for GPU acceleration packages. Not applicable on systems without NVIDIA GPUs or on unsupported platforms.
Install
jax-cuda12-plugin on PyPI
Before you install
Medium install friction due to platform-specific wheels for Linux x86_64 and aarch64 only, plus a compiled runtime dependency on jax-cuda12-pjrt. Actively maintained with recent releases and high repository activity (36158 stars).
Requires NVIDIA GPU with CUDA 12 support and Linux x86_64 or aarch64 platform; Python 3.12 or later.
License in practice
Apache-2.0 permissive license allows commercial and private use with minimal restrictions.
Quickstart
pip install jax-cuda12-plugin
import jax
import jax.numpy as jnp
# GPU acceleration is automatically available for JAX operations
result = jax.jit(lambda x: jnp.dot(x, x))(jnp.ones((100, 100)))
Verify before relying
- Whether CUDA 12 runtime libraries must be pre-installed on the system for the plugin to function.
- Performance characteristics and memory overhead compared to CPU-only JAX execution.
- Compatibility with specific NVIDIA GPU architectures and CUDA driver versions.
Package facts
| License | Apache-2.0 permissive |
| Python support | Supports the current Python release >=3.12 |
| Install friction | Medium. Platform-specific wheel |
| Runtime dependencies | 1 packagejax-cuda12-pjrt |
| Maintenance | Actively maintained 29 days since the last release |
| Last repo commit | |
| First released | |
| Downloads | 782,417 / month, #5,078 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 :: 3.12Programming Language :: Python :: 3.13Programming Language :: Python :: 3.14Programming Language :: Python :: Free Threading :: 3 - Stable |
Evidence: jax_cuda12_plugin-0.11.0-cp312-cp312-manylinux_2_27_aarch64.whl; jax_cuda12_plugin-0.11.0-cp312-cp312-manylinux_2_27_x86_64.whl; jax_cuda12_plugin-0.11.0-cp313-cp313-manylinux_2_27_aarch64.whl; jax_cuda12_plugin-0.11.0-cp313-cp313-manylinux_2_27_x86_64.whl; jax_cuda12_plugin-0.11.0-cp314-cp314-manylinux_2_27_aarch64.whl; jax_cuda12_plugin-0.11.0-cp314-cp314-manylinux_2_27_x86_64.whl; jax_cuda12_plugin-0.11.0-cp314-cp314t-manylinux_2_27_aarch64.whl; jax_cuda12_plugin-0.11.0-cp314-cp314t-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 gpu backend”
- jax-cuda12-pluginEnables JAX to run numerical computations and machine learning…
- jax-cuda13-pluginProvides NVIDIA GPU support for JAX by enabling CUDA 13 compilation…
- jax-cuda12-pjrtProvides NVIDIA GPU acceleration for JAX numerical computing via the…
Give your agent the search over MCP, or paste the wish link into any chat.
More Artificial Intelligence packages
LiteLLM provides a unified Python interface to call 100+ LLM providers (OpenAI, Anthropic, Gemini, Bedrock, Azure, and others) using OpenAI-compatible API format, available as both a Python SDK and a self-hosted AI Gateway proxy server.
Install it if you need to work with multiple LLM providers or want to centralize LLM routing in your organization.
Client library and CLI tool for downloading, uploading, and managing models, datasets, and repositories on the Hugging Face Hub platform.
Install it if you work with Hugging Face Hub models or datasets.
LangChain provides a framework for building agents and LLM-powered applications by composing language models, tools, and memory through a unified API that abstracts over multiple model providers.
hf-xet provides chunk-based deduplication and efficient file transfer for the Hugging Face Hub, enabling faster uploads and downloads of large files with local disk caching.
Tokenizers converts raw text into token sequences for NLP models, with support for training custom vocabularies and using pre-built tokenizers (BPE, WordPiece) optimized for speed via Rust.
Transformers provides a unified framework for loading, fine-tuning, and running state-of-the-art pretrained models across text, vision, audio, video, and multimodal tasks using PyTorch, JAX, or TensorFlow.
Install it if you need to run or train any transformer-based model for NLP, vision, audio, or multimodal tasks.
See also augmax · jax · jax-cuda13-plugin · jaxlib · torchax · tokamax · jax-cuda13-pjrt · jax-cuda12-pjrt · jax-jumpy · jmp