Anyone from the Numba team care to comment?
JAX core is an extensible system for transforming numerical Python functions. This core is used to implement automatic differentiation, translation to TF XLA, etc.
Numba does not have such a generic function transformation framework - it just supports a single transformation, that is from numerical Python function to machine code.
We haven't tried combining them yet, but we think it would be fun to explore (https://github.com/google/jax/issues/1870). For example, you could use Numba to hand write a numerical kernel that then participates in a machine learning model that uses JAX automatic differentiation.
(I'm not from either jax or numba, but a keen jax user for non-ML research.)