Chapter 14¶
Overview¶
Profiling and optimizing JAX code for production systems.
Topics¶
1. Profiling JAX Code¶
import jax
import jax.numpy as jnp
import time
def benchmark(f, x, n_runs=100):
# Warmup
for _ in range(10):
f(x)
# Time
start = time.time()
for _ in range(n_runs):
f(x).block_until_ready()
elapsed = time.time() - start
return elapsed / n_runs
# Define functions
def f_no_jit(x):
return jnp.sin(x) + jnp.cos(x)**2
f_jit = jax.jit(f_no_jit)
x = jnp.arange(1e6)
t_no_jit = benchmark(f_no_jit, x)
t_jit = benchmark(f_jit, x)
print(f"No JIT: {t_no_jit*1000:.2f}ms")
print(f"JIT: {t_jit*1000:.2f}ms")
print(f"Speedup: {t_no_jit/t_jit:.1f}x")
2. Memory Profiling¶
import tracemalloc
def profile_memory(f, x):
tracemalloc.start()
result = f(x)
result.block_until_ready()
current, peak = tracemalloc.get_traced_memory()
tracemalloc.stop()
return current / 1e6, peak / 1e6
current_mb, peak_mb = profile_memory(f_jit, x)
print(f"Current: {current_mb:.1f}MB, Peak: {peak_mb:.1f}MB")
3. Identifying Bottlenecks¶
# Use jax.debug.print for debugging
@jax.jit
def f_debug(x):
jax.debug.print("x shape: {}", x.shape)
y = jnp.sin(x)
jax.debug.print("y shape: {}", y.shape)
return y
# Or disable JIT temporarily
with jax.disable_jit():
result = f(x)
Summary¶
- Profile before optimizing
- Use block_until_ready() for timing
- Memory profiling with tracemalloc
- Debug with jax.debug.print
- JIT compilation is key optimization