@jax.jit def suspicious_function(x, w): hidden = jnp.dot(x, w) jax.debug.breakpoint() # Execution pauses here at RUNTIME return jax.nn.relu(hidden) output = suspicious_function(x, w) __ __