import jax import jax.numpy as jnp from jax import lax # Runtime Printing Inside JIT jax.debug.print("value = {}", x) jax.debug.print("step {} loss {}", step, loss, ordered=True) # Interactive Breakpoint Inside JIT jax.debug.breakpoint() # Conditional breakpoint lax.cond(has_nan, lambda: jax.debug.breakpoint(), lambda: None) # Full Eager Execution (standard pdb/print work) with jax.disable_jit(): result = my_function(x) jax.config.update("jax_disable_jit", True) # Global toggle # Automatic NaN Detection jax.config.update("jax_debug_nans", True) # Flax NNX Structural Inspection nnx.display(model) # Capturing Intermediate Values self.sow(nnx.Intermediate, 'my_activation', x) model.my_activation.value # Access after forward pass __ __