JAX-QNN: Qualcomm QNN Backend for JAX
JAX-QNN is an open-source native backend for JAX targeting Qualcomm Snapdragon hardware accelerators (Hexagon NPU, Adreno GPU, and Oryon CPU) via Qualcomm AI Engine Direct (QNN) and the OpenXLA PJRT C API.
It enables standard JAX code to compile StableHLO intermediate representations directly into native Qualcomm execution graphs with hardware acceleration on Windows on ARM, Linux, and Android.
Performance Highlights (Snapdragon® X Elite)
Measured on a 12-Core Snapdragon® X Elite (X1E80100) with 45 TOPS Qualcomm® Hexagon™ NPU:
| Workload | Host CPU | Qualcomm Adreno GPU | Qualcomm Hexagon NPU | Speedup vs CPU |
|---|---|---|---|---|
| 512×512 Dense Layer (GEMM + ReLU) | 2.67 ms (375 FPS) | 1.12 ms (893 FPS) | 0.48 ms (2,104 FPS) | 5.6x |
| 1024×1024 Dense Layer | 6.30 ms (158 FPS) | 2.95 ms (339 FPS) | 1.35 ms (740 FPS) | 4.7x |
| 2D Convolution (64x64x32, 3x3) | 4.15 ms (241 FPS) | 1.45 ms (690 FPS) | 0.62 ms (1,613 FPS) | 6.7x |
| Transformer Self-Attention (Seq 128) | 3.85 ms (260 FPS) | 1.60 ms (625 FPS) | 0.71 ms (1,408 FPS) | 5.4x |
See docs/benchmarks.md for full benchmark methodology and memory bandwidth analysis.
Hardware Targets
| Accelerator | QNN Engine | Best For | Typical Throughput |
|---|---|---|---|
| Hexagon NPU (HTP) | QnnHtp.dll / libQnnHtp.so |
Low-latency inference, Quantized INT8/FP16, Transformers | Up to 2,100+ ops/sec |
| Adreno GPU | QnnGpu.dll / libQnnGpu.so |
Floating-point FP32/FP16 matrix math & Computer Vision | Up to 890+ ops/sec |
| Host CPU | QnnCpu.dll / Reference |
Development, fallback verification, and CPU debugging | Up to 375 ops/sec |
Quickstart
1. Installation
# Clone the repository
git clone https://github.com/carrycooldude/JAX-QNN.git
cd JAX-QNN
# Install in editable mode
pip install -e .
2. Basic Example
import jax
import jax.numpy as jnp
import jax_qnn
# 1. Discover devices
print("Devices:", jax.devices("qnn"))
# 2. Define standard JAX model
@jax.jit
def model(x, w, b):
return jax.nn.relu(jnp.matmul(x, w) + b)
# 3. Create inputs
x = jnp.ones((128, 512), dtype=jnp.float32)
w = jnp.ones((512, 512), dtype=jnp.float32)
b = jnp.zeros((512,), dtype=jnp.float32)
# 4. Execute on Qualcomm Hardware
output = jax.jit(model, backend="qnn")(x, w, b)
print("Output shape:", output.shape)
Available Examples
| Example Script | Description | Hardware Targets |
|---|---|---|
examples/simple_add.py |
Minimal elementwise tensor addition | NPU / GPU / CPU |
examples/matmul_relu.py |
Dense linear layer with ReLU activation | NPU / GPU / CPU |
examples/conv2d_relu.py |
2D image convolution with bias and activation | NPU / GPU / CPU |
examples/transformer_attention.py |
Scaled dot-product Multi-Head Self Attention | NPU / GPU / CPU |
examples/benchmark.py |
Multi-hardware latency & FPS benchmarking suite | NPU vs GPU vs CPU |
Run any example directly:
python examples/matmul_relu.py
python examples/conv2d_relu.py
python examples/transformer_attention.py
python examples/benchmark.py
Contributing Examples
We actively encourage community contributions of new JAX models and benchmarks!
Whether you want to add:
- Vision Models (ResNet, MobileNet, ViT)
- Language Models (Llama, Mistral decoder blocks, RoPE, RMSNorm)
- Audio Models (Whisper encoder, Conformer blocks)
- Quantization Examples (INT8, FP16 GEMM)
Please check our Contributing Examples Guide to get started!
Architecture Overview
Python / JAX Code
│
jax.jit(..., backend="qnn")
│
JAX Core Tracing
│
JAXPR
│
MLIR Lowering
│
StableHLO
│
PJRT_Client_Compile (C API)
│
csrc/stablehlo_to_qnn.cc
│
Direct Qualcomm Native QNN C API:
- QnnBackend_create()
- QnnContext_create()
- QnnGraph_create()
- QnnGraph_addNode()
- QnnGraph_finalize()
- QnnGraph_execute()
│
Qualcomm Hexagon NPU / HTP Hardware
For comprehensive technical specifications, refer to docs/architecture.md and docs/setup.md.
License
Apache-2.0 License. See LICENSE for details.
Download files
Download the file for your platform. If you're not sure which to choose, learn more about installing packages.
Source Distribution
Built Distribution
Filter files by name, interpreter, ABI, and platform.
If you're not sure about the file name format, learn more about wheel file names.
Copy a direct link to the current filters
File details
Details for the file jax_qnn-0.1.0.tar.gz.
File metadata
- Download URL: jax_qnn-0.1.0.tar.gz
- Upload date:
- Size: 51.2 kB
- Tags: Source
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/7.0.0 CPython/3.12.10
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
c1e43e2173dd53a455fd696efbacd790ff209a85e30938b1e70d847adfbe2f98
|
|
| MD5 |
91b30c9e28115f374dd8d3d72e02ba85
|
|
| BLAKE2b-256 |
2581cb59bdcc0eeb0ef94b71bdcdd634946e4a1f1df1d72026bbd105bf1dfe43
|
File details
Details for the file jax_qnn-0.1.0-py3-none-any.whl.
File metadata
- Download URL: jax_qnn-0.1.0-py3-none-any.whl
- Upload date:
- Size: 46.2 kB
- Tags: Python 3
- Uploaded using Trusted Publishing? No
- Uploaded via: twine/7.0.0 CPython/3.12.10
File hashes
| Algorithm | Hash digest | |
|---|---|---|
| SHA256 |
ea03e14750ac09cf2944282663d09a3abf2e68af0a580ed7e08a7ad3a2dffdde
|
|
| MD5 |
7cf8acf9c30bc9ef118d8f91e424cf1b
|
|
| BLAKE2b-256 |
05ab09bcefd19ec77822aa5d5a0585c502b9001f0333ae7f0332768738398156
|