import jax jax.config.update("jax_debug_nans", True) @jax.jit def buggy_function(x): y = jnp.log(x) # log(0) = -inf, log(negative) = NaN return y * 2 # This will halt immediately with a traceback pointing at jnp.log result = buggy_function(jnp.array([-1.0, 2.0, 3.0])) __ __