JAX Under the Hood: Debugging and Performance Profiling using XProf

About this session

JAX is the engine behind modern LLM training, offering incredible speed through XLA compilation. However, its functional paradigm, tracing behavior, and asynchronous execution require a major mental shift. How do you master these concepts and learn to debug them when time is short?

In this guided lab, I will demystify the JAX execution model. Designed to be highly efficient, this session utilizes a pre-configured "Broken vs. Fixed" Google Colab notebook. Instead of writing boilerplate, attendees will immediately interact with live code to witness and fix real world engineering pitfalls.

I will cover the "JAX Trinity" (jit, vmap, grad), and then dive into two high-impact debugging scenarios: combating shape mismatches/recompilation, and resolving Out-of-Memory (OOM) errors using rematerialization (remat). Finally, I will show how to analyze these workloads using the XProf in Tensorboard

You will leave with a solid mental model of JAX and a ready-to-use debugging notebook for your own production pipelines.

Speaker

Key takeaways

  • Learn to use jax.debug primitives to inspect device memory values without disrupting compilation graphs.
  • Isolate shape mismatches and implement rematerialization (remat) to resolve tracking errors and Out-of-Memory faults.
  • Utilize XProf to evaluate HLO graphs, optimize Model Flops Utilization (MFU), and audit device memory footprints.

Related sessions