winjax
Native Windows CUDA support for JAX — no WSL2 required. Unofficial.
winjax installs a PJRT GPU plugin built
natively for Windows from XLA sources, plus a small loader that registers it
with stock jax/jaxlib. All CUDA runtime libraries (CUDA 13, cuDNN 9) come
from NVIDIA's pip wheels, so a working NVIDIA driver is the only system
requirement.
Requirements
- Windows 10/11, x86-64
- Python 3.13
- An NVIDIA GPU with a driver supporting CUDA 13
Install
pip install winjax
Use
import jax
print(jax.devices()) # [CudaDevice(id=0)]
Nothing else to configure: import jax discovers the plugin through the
jax_plugins namespace package.
Packages
winjax— loader (this package, pure Python)winjax-cuda13-pjrt— the Windows-built XLA CUDA PJRT plugin DLLwinjax-cuda13-plugin— CUDA kernel extension modules (jax_cuda13_plugin)
Source and patches: https://github.com/eterevsky/winjax
Metadata
Release files for winjax 0.11.2
For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.
Built distribution (wheel)
| File | Interpreter | ABI | Platform | Reset |
|---|---|---|---|---|
| winjax-0.11.2-py3-none-any.whl | Python 3 | none | any | Details |
Release files / winjax-0.11.2-py3-none-any.whl
| Download URL | winjax-0.11.2-py3-none-any.whl |
|---|---|
| Size | 8.2 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
16497cc291cd744f4d10b187bf4723cf1c6237615c1dfe16209a11335771220a
|
|
BLAKE2b-256 checksum How to use checksums |
1a9141e7ea360277dc6c47f656c7d36f823c328d9b0cac6c9a66f4e9dac6bd62
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
uv/0.12.5 {"installer":{"name":"uv","version":"0.12.5","subcommand":["publish"]},"python":null,"implementation":{"name":null,"version":null},"distro":null,"system":{"name":null,"release":null},"cpu":null,"openssl_version":null,"setuptools_version":null,"rustc_version":null,"ci":null}
|