$npx skillfedfor your agent

jax-cuda12-plugin

JAX Plugin for NVIDIA GPUs

With conditionsPyPI Artificial IntelligenceReleased Jul 2026782.4K downloads / moApache-2.0Platform wheel

Decision gist · record as of 2026-08-14

platform wheels — 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
v0.11.0 · released 2026-07-16 · Python >=3.12 · 1 runtime deps: jax-cuda12-pjrt

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

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.
Same gist for agents: .md · .json

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.

With conditions

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

LicenseApache-2.0 permissive
Python supportSupports the current Python release >=3.12
Install frictionMedium. Platform-specific wheel
Runtime dependencies
1 package
jax-cuda12-pjrt
MaintenanceActively maintained 29 days since the last release
Last repo commit
First released
Downloads782,417 / month, #5,078 on PyPI 30-day window, as of 2026-08-14
Known vulnerabilitiesNone 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

Capabilities
jax nvidia gpu accelerationcuda 12 gpu pluginjax gpu backendnvidia cuda jax supportgpu-accelerated numerical computing
Topics
gpu-accelerationcudanumerical-computing

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 With conditions
PyPI · Artificial Intelligence · released Aug 2026

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.

MITcompiled wheel
682.8Mdownloads / mo
huggingface-hub Worth it
PyPI · Artificial Intelligence · released Aug 2026

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.

Apache-2.0pure Python · 3.10.0+
442.4Mdownloads / mo
langchain Worth it
PyPI · Python Modules · released Aug 2026

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.

MITpure Python
315.4Mdownloads / mo
hf-xet With conditions
PyPI · Artificial Intelligence · released Aug 2026

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.

Apache-2.0compiled wheel · 3.8+
258.4Mdownloads / mo
tokenizers Worth it
PyPI · Artificial Intelligence · released Apr 2026

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.

Apache-2.0compiled wheel · 3.10+
222.9Mdownloads / mo
transformers Worth it
PyPI · Artificial Intelligence · released Aug 2026

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.

permissive licensepure Python · 3.10.0+
186.6Mdownloads / mo

See also augmax · jax · jax-cuda13-plugin · jaxlib · torchax · tokamax · jax-cuda13-pjrt · jax-cuda12-pjrt · jax-jumpy · jmp

Further reading