Complete, copy-pasteable JAX models ready to compile and execute on Qualcomm Snapdragon hardware.
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)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)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)