import jax import jax.numpy as jnp @jax.jit def debug_forward(x, w): logits = jnp.dot(x, w) jax.debug.print("logits = {}", logits) return jax.nn.softmax(logits) x = jnp.array([[1.0, 2.0, 3.0]]) w = jnp.array([[0.1, 0.2], [0.3, 0.4], [0.5, 0.6]]) output = debug_forward(x, w) __ __