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.
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
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_ |
the k largest values of the last axis and their indices, as jax.lax.top_ |
rank(x, axis=-1, method="average", descending=False) |
each value's rank, as scipy.stats.rankdata |
Notes
- They follow JAX's rules:
-0.0and+0.0tie, and NaNs sort last. - They run on the CPU backend; arrays on other devices stay with JAX.
Related
- Sorting operators and layers: the same calls for PyTorch.
- Python API reference