Skip to content
← All tracks

JAX Vectorization & Control Flow

JAX

vmap, scan, while_loop, fori_loop, cond, switch — JAX's primitives for batching and control flow.

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

0 / 25 solved
  1. 1. Not solved yet. vmap with in_axes
  2. 2. Not solved yet. vmap with out_axes
  3. 3. Not solved yet. vmap with None broadcasting
  4. 4. Not solved yet. Nested vmap (vmap-of-vmap)
  5. 5. Not solved yet. Cumulative Sum via lax.scan
  6. 6. Not solved yet. Running Mean via lax.scan
  7. 7. Not solved yet. Scan over Layer Stack
  8. 8. Not solved yet. Training Loop via lax.scan
  9. 9. Not solved yet. Scan with Per-Step Outputs
  10. 10. Not solved yet. Newton's Method via lax.while_loop
  11. 11. Not solved yet. Bounded Search via lax.while_loop
  12. 12. Not solved yet. lax.while_loop vs Python while
  13. 13. Not solved yet. x^n via lax.fori_loop
  14. 14. Not solved yet. fori vs scan vs while
  15. 15. Not solved yet. Binary Classification via lax.cond
  16. 16. Not solved yet. Multi-Branch Dispatch via lax.switch
  17. 17. Not solved yet. vmap over lax.cond
  18. 18. Not solved yet. vmap over lax.scan
  19. 19. Not solved yet. jit of scan
  20. 20. Not solved yet. vmap-cond vs where: cost tradeoff
  21. 21. Not solved yet. Dynamic Slice
  22. 22. Not solved yet. Dynamic Update Slice
  23. 23. Not solved yet. Associative Scan (Parallel Cumsum)
  24. 24. Not solved yet. lax.scan with reverse=True
  25. 25. Not solved yet. lax.map vs vmap

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

What does lax.scan return, and why use it over a Python loop?

def step(c, x):
    return c + x, c

carry, ys = jax.lax.scan(step, 0.0, jnp.arange(5.0))
Question 1 of 4