import jax def my_function(x, w): hidden = jnp.dot(x, w) print(f"hidden = {hidden}") # Works normally now! breakpoint() # Standard Python debugger works! return jax.nn.relu(hidden) with jax.disable_jit(): output = my_function(x, w) __ __