2 ms·
JAX recompiles functions every time you call them with an array of a new shape, i.e. if called with an array of shape (7,10) and then one with shape (17, 12) th
by adament 5y ago
JAX recompiles functions every time you call them with an array of a new shape, i.e. if called with an array of shape (7,10) and then one with shape (17, 12) the function is compiled twice. This is generally fine for deep learning and some other numerical applications where you do the same computations again and again over arrays with the same shape. But for data exploration like Pandas, in my experience, your data shapes are different with each call, so the repeated recompilations make it unattractive.
In Numba it only recompiles for dimension changes i.e. if shape changes from (7, 10) to (2, 10, 12). However numba does not integrate AD so if you need that, JAX is probably your best bet unless Enzyme matures and you want to look into integratig it with numba.