Runnable Examples & Workloads

Complete, copy-pasteable JAX models ready to compile and execute on Qualcomm Snapdragon hardware.

1. Dense Layer: MatMul + ReLU

A classic fully-connected layer running via backend="qnn":

import jax
import jax.numpy as jnp
import jax_qnn

@jax.jit
def dense_layer(x, w, b):
    return jax.nn.relu(jnp.matmul(x, w) + b)

x = jnp.ones((128, 512), dtype=jnp.float32)
w = jnp.ones((512, 512), dtype=jnp.float32)
b = jnp.zeros((512,), dtype=jnp.float32)

y = jax.jit(dense_layer, backend="qnn")(x, w, b)
print("Output:", y.shape)

2. 2D Convolution (CNN Vision Block)

Image convolution with 3x3 kernel and bias addition:

import jax
import jax.numpy as jnp
import jax.lax as lax
import jax_qnn

@jax.jit
def conv_block(x, kernel, bias):
    conv_out = lax.conv_general_dilated(
        lhs=x, rhs=kernel, window_strides=(1, 1),
        padding='SAME', dimension_numbers=('NHWC', 'HWIO', 'NHWC')
    )
    return jax.nn.relu(conv_out + bias)

x = jnp.ones((1, 64, 64, 32), dtype=jnp.float32)
kernel = jnp.ones((3, 3, 32, 64), dtype=jnp.float32) / 9.0
bias = jnp.zeros((64,), dtype=jnp.float32)

out = jax.jit(conv_block, backend="qnn")(x, kernel, bias)
print("Output Shape:", out.shape)

3. Transformer Multi-Head Self-Attention

Scaled dot-product attention block targeting Hexagon NPU Vector Cores:

import jax
import jax.numpy as jnp
import jax_qnn

@jax.jit
def self_attention(q, k, v):
    d_k = q.shape[-1]
    scores = jnp.matmul(q, jnp.swapaxes(k, -2, -1)) / jnp.sqrt(d_k)
    weights = jax.nn.softmax(scores, axis=-1)
    return jnp.matmul(weights, v)

q = jnp.ones((1, 4, 128, 64), dtype=jnp.float32)
k = jnp.ones((1, 4, 128, 64), dtype=jnp.float32)
v = jnp.ones((1, 4, 128, 64), dtype=jnp.float32)

out = jax.jit(self_attention, backend="qnn")(q, k, v)
print("Attention Out:", out.shape)