Chapter 16¶
Overview¶
Debugging functional, JIT-compiled code requires different techniques than standard Python.
Topics¶
1. Common Errors¶
Error: Non-hashable pytrees
# Wrong
def f(params): # params is a list
return jnp.sum(jnp.array(params)**2)
# Right
def f(params): # params is a tuple
return jnp.sum(jnp.array(params)**2)
Error: Invalid shapes with JIT
# Wrong
@jax.jit
def f(x):
if x.shape[0] > 5:
return x[:5]
else:
return x
# Right
@jax.jit
def f(x):
return lax.cond(
x.shape[0] > 5,
lambda: x[:5],
lambda: x
)
2. Debug Techniques¶
Disable JIT
with jax.disable_jit():
# Run without JIT to see actual errors
result = f(x)
Use jax.debug.print
@jax.jit
def f(x):
jax.debug.print("x value: {}", x)
return x**2
Numerical Gradient Checking
def check_gradient(f, x, eps=1e-5):
# Numerical
grad_num = (f(x + eps) - f(x - eps)) / (2 * eps)
# Autodiff
grad_auto = jax.grad(f)(x)
error = jnp.abs(grad_num - grad_auto)
print(f"Gradient error: {error}")
assert error < 1e-4, "Gradients don't match!"
3. Tracing Issues¶
Understanding Tracing
import jax
import jax.numpy as jnp
@jax.jit
def f(x):
# This gets traced with an abstract value
# You can't use x.shape in conditionals
print(f"x during tracing: {x}") # Prints abstract value
return x**2
f(jnp.array([1., 2., 3.])) # Traces and compiles
f(jnp.array([1., 2., 3., 4.])) # Recompiles (different shape)
Summary¶
- Disable JIT to debug
- Use jax.debug.print for diagnostics
- Check gradients numerically
- Understand tracing with abstract values
- Common errors: non-hashable pytrees, shape dependencies