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.0.post6
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.0.post6-py3-none-any.whl | Python 3 | none | any | Details |
Release files / winjax-0.11.0.post6-py3-none-any.whl
| Download URL | winjax-0.11.0.post6-py3-none-any.whl |
|---|---|
| Size | 8.2 kB |
| Tags | Python 3 |
|
SHA-256 checksum How to use checksums |
777eed7a649a2b03cbf53b96f2fc905fbd458cab8f3a74a94d060fcc9b5e2d89
|
|
BLAKE2b-256 checksum How to use checksums |
2d34a3d52308f8e5ddce6e5abc6fd0ae765c8cecf7fb59c0f3f9a05ca19ea13e
|
| Upload date | |
|
Uploaded using Trusted Publishing? What is trusted publishing? |
No |
| Uploaded via |
uv/0.11.8 {"installer":{"name":"uv","version":"0.11.8","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}
|