@jax.jit def broken_check(x): # This will crash with a ConcretizationTypeError assert jnp.all(x > 0), "Values must be positive!" return x * 2 __ __