Metadata-Version: 2.4
Name: winjax
Version: 0.11.0
Summary: Native Windows CUDA support for JAX (unofficial PJRT plugin loader)
Author-email: Oleg Eterevsky <oleg@eterevsky.com>
License-Expression: Apache-2.0
Project-URL: Homepage, https://github.com/eterevsky/winjax
Classifier: Development Status :: 3 - Alpha
Classifier: Environment :: GPU :: NVIDIA CUDA :: 13
Classifier: Operating System :: Microsoft :: Windows
Classifier: Programming Language :: Python :: 3.13
Classifier: Topic :: Scientific/Engineering
Requires-Python: <3.15,>=3.13
Description-Content-Type: text/markdown
Requires-Dist: jax==0.11.0
Requires-Dist: jaxlib==0.11.0
Requires-Dist: winjax-cuda13-pjrt==0.11.0.*
Requires-Dist: winjax-cuda13-plugin==0.11.0.*
Requires-Dist: nvidia-cublas<14,>=13.0.0.19
Requires-Dist: nvidia-cuda-cupti<14,>=13.3.75
Requires-Dist: nvidia-cuda-nvcc<14,>=13.0.48
Requires-Dist: nvidia-cuda-nvrtc<14,>=13.0.48
Requires-Dist: nvidia-cuda-runtime<14,>=13.0.48
Requires-Dist: nvidia-cudnn-cu13<10,>=9.24.0.43
Requires-Dist: nvidia-cufft<13,>=12.0.0.15
Requires-Dist: nvidia-cusolver<13,>=12.0.3.29
Requires-Dist: nvidia-cusparse<13,>=12.6.2.49
Requires-Dist: nvidia-nvjitlink<14,>=13.0.39
Requires-Dist: nvidia-nvvm<14,>=13

# winjax

Native Windows CUDA support for JAX — no WSL2 required. Unofficial.

`winjax` installs a [PJRT](https://openxla.org/xla/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

```python
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 DLL
- `winjax-cuda13-plugin` — CUDA kernel extension modules (`jax_cuda13_plugin`)

Source and patches: https://github.com/eterevsky/winjax
