import jax.numpy as jnp from jax import lax @jax.jit def conditional_breakpoint(x): has_nan = jnp.any(jnp.isnan(x)) lax.cond( has_nan, lambda: jax.debug.breakpoint(), lambda: None, ) return x * 2 __ __