@jax.jit def custom_softmax(x): probs = jax.nn.softmax(x) # This crashes with ConcretizationTypeError assert jnp.all(probs >= 0), "Probabilities cannot be negative!" return probs __ __