Kwker

Accelerate JAX

kwker.jax_ops gives JAX four calls that run as XLA custom calls on the CPU: sort, argsort, top_k and rank. They work under jit, vmap and differentiation, with no copy to the host. A float32 sort of a million keys ran 86.7x faster than jnp.sort.

PythonNeeds jax: runs on your machine.
import jax
import jax.numpy as jnp
import kwker.jax_ops as sj

x = jnp.array([3.0, 1.0, 2.0])
print(sj.sort(x), sj.argsort(x))
Output
[1. 2. 3.] [1 2 0]

Inside jit and vmap

PythonNeeds jax: runs on your machine.
import jax
import jax.numpy as jnp
import kwker.jax_ops as sj

@jax.jit
def best(scores):
    return sj.top_k(scores, 2)

batch = jnp.array([[0.1, 0.9, 0.5], [0.7, 0.2, 0.8]])
values, indices = jax.vmap(best)(batch)
print(indices)
Output
[[1 2]
 [2 0]]

The calls

Call Result
sort(x, axis=-1, descending=False) the sorted array
argsort(x, axis=-1, descending=False) the indices that sort x, stable
top_k(x, k) the k largest values of the last axis and their indices, as jax.lax.top_k
rank(x, axis=-1, method="average", descending=False) each value's rank, as scipy.stats.rankdata

Notes