Skip to content
← All tracks

JAX Profiling & Debugging

JAX

Profiler hooks, debug primitives, fixing common tracer/concretization errors, AOT compilation. Practical-skill problems for production work.

This track is written for JAX, which isn't the mode you're browsing in.

0 / 12 solved
  1. 1. Not solved yet. jax.named_scope for Profiler Grouping
  2. 2. Not solved yet. profiler.StepTraceAnnotation
  3. 3. Not solved yet. @profiler.annotate_function
  4. 4. Not solved yet. jax.debug.print with Multiple kwargs
  5. 5. Not solved yet. jax.debug.callback for Side Effects
  6. 6. Not solved yet. jnp.allclose for Tolerance Equality
  7. 7. Not solved yet. Fix: List Append → jnp.concatenate
  8. 8. Not solved yet. Fix: Python if → jnp.where
  9. 9. Not solved yet. Fix: Python for → vmap
  10. 10. Not solved yet. jaxpr with Multiple Args
  11. 11. Not solved yet. AOT via jit.lower().compile()
  12. 12. Not solved yet. Pretty-Printed jaxpr Length

Check yourself

3 questions · one attempt each

These do not count toward finishing the track. They are here to catch the things that are easy to read past.

0 / 3

You are timing a jitted function. Which two things must you do for the number to mean anything?

# t0 = time(); out = f(x); t1 = time()
Question 1 of 3