Skip to content
← All tracks

JAX Foundations

JAX

Functional ML — pure functions, vmap, pytrees. The JAX way of thinking.

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

0 / 25 solved
  1. 1. Not solved yet. Vectorize with vmap
  2. 2. Not solved yet. Pure Function vs Impure
  3. 3. Not solved yet. Tracing: Shape vs Value
  4. 4. Not solved yet. JIT Compile a Function
  5. 5. Not solved yet. Jit with static_argnames
  6. 6. Not solved yet. Tracer Leak Detection
  7. 7. Not solved yet. Functional Update with .at[].set()
  8. 8. Not solved yet. Functional Update with .at[].add()
  9. 9. Not solved yet. 2-D Scatter via .at[].add()
  10. 10. Not solved yet. Gradient with jax.grad
  11. 11. Not solved yet. jax.value_and_grad
  12. 12. Not solved yet. jit + grad Composition
  13. 13. Not solved yet. PRNGKey and Split
  14. 14. Not solved yet. PRNGKey fold_in
  15. 15. Not solved yet. Deterministic Batch via vmap+split
  16. 16. Not solved yet. Pytree Leaves
  17. 17. Not solved yet. Pytree Map
  18. 18. Not solved yet. Gradient over Pytree Params
  19. 19. Not solved yet. JAX numpy vs PyTorch ops
  20. 20. Not solved yet. Dtype Promotion Rules
  21. 21. Not solved yet. take_along_axis (Gather)
  22. 22. Not solved yet. Multi-Condition where
  23. 23. Not solved yet. NaN-Safe Mean (with mask)
  24. 24. Not solved yet. Explicit Broadcasting via broadcast_to
  25. 25. Not solved yet. Clip and Extrema

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

The function is called three times: twice with a scalar, then once with a shape (2,) array. How many times does the Python print run?

@jax.jit
def f(x):
    print("tracing, x =", x)
    return x * 2

f(jnp.array(1.0))
f(jnp.array(3.0))
f(jnp.ones((2,)))
Question 1 of 4