Skip to content
← All tracks

JAX Autodiff

JAX

Forward and reverse mode, custom derivatives, gradient checkpointing, per-example gradients. The autodiff toolbox.

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

0 / 25 solved
  1. 1. Not solved yet. jvp Basics
  2. 2. Not solved yet. jacfwd vs jacrev
  3. 3. Not solved yet. jvp for Sensitivity Analysis
  4. 4. Not solved yet. vjp Basics
  5. 5. Not solved yet. Jacobian via Batched vjp
  6. 6. Not solved yet. grad vs vjp
  7. 7. Not solved yet. Hessian of a Quadratic
  8. 8. Not solved yet. HVP via grad-of-grad
  9. 9. Not solved yet. HVP via jvp-of-grad
  10. 10. Not solved yet. Custom VJP: Stable log1pexp
  11. 11. Not solved yet. Custom JVP: Clip with Pass-through Gradient
  12. 12. Not solved yet. Custom VJP: Implicit Function Theorem
  13. 13. Not solved yet. stop_gradient: Target Network
  14. 14. Not solved yet. Straight-Through Estimator
  15. 15. Not solved yet. Gradient Checkpointing: Basics
  16. 16. Not solved yet. Checkpoint with Save Policy
  17. 17. Not solved yet. Checkpointed Deep Stack via scan
  18. 18. Not solved yet. Per-Example Gradients via vmap(grad(...))
  19. 19. Not solved yet. vmap(grad) vs grad(sum(vmap))
  20. 20. Not solved yet. Microbatched Gradient Accumulation via scan
  21. 21. Not solved yet. jax.linearize Primitive
  22. 22. Not solved yet. Jacobian via Mixed-Mode (jvp+vjp)
  23. 23. Not solved yet. Higher-Order custom_vjp
  24. 24. Not solved yet. Saved Residuals in custom_vjp
  25. 25. Not solved yet. Grad through stop_gradient

Check yourself

4 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 / 4

You need both the loss and its gradient each step. What does value_and_grad save?

jax.value_and_grad(lambda x: (x ** 2).sum())(jnp.array([3.0]))
Question 1 of 4