JAX 401: kernels and FFI

JAX 401: kernels and FFI#

Sometimes the compiler isn’t enough: you want to write the device code yourself, or call code that lives outside JAX entirely. This section covers both escape hatches.

  1. Pallas: custom kernels in JAX — Pallas, JAX’s kernel language for writing custom GPU and TPU kernels; mostly a map to the extensive Pallas documentation at pallas.jax.dev.

  2. Writing High-Performance GPU Kernels with CuTe DSL and JAX — writing high-performance GPU kernels with NVIDIA’s CuTe DSL and calling them from JAX.

  3. Foreign function interface (FFI) — the foreign function interface: wrapping external C++/CUDA code as a JAX operation, and teaching it to work with jit, vmap, grad, and sharding.