Skip to content
← All tracks

JAX Conceptual Deep-Dives

JAX

Custom pytree nodes, host interop callbacks, quantization patterns, jaxpr inspection, advanced debugging. Beyond the curriculum.

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

0 / 15 solved
  1. 1. Not solved yet. register_pytree_node
  2. 2. Not solved yet. tree_flatten / tree_unflatten Round-Trip
  3. 3. Not solved yet. Pytree Aux Data
  4. 4. Not solved yet. pure_callback Basics
  5. 5. Not solved yet. io_callback for Side Effects
  6. 6. Not solved yet. jax.debug.print
  7. 7. Not solved yet. JAX <-> NumPy Bridge
  8. 8. Not solved yet. int8 Affine Quantization
  9. 9. Not solved yet. int8 Dequantization
  10. 10. Not solved yet. float8 Round-Trip
  11. 11. Not solved yet. Counting Jaxpr Equations
  12. 12. Not solved yet. jax.debug.callback
  13. 13. Not solved yet. Selective stop_gradient (Frozen Layers)
  14. 14. Not solved yet. jvp with Structured Tangents
  15. 15. Not solved yet. jaxpr with jit

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 jax.tree.map do with a nested dict of arrays?

tree = {"a": jnp.array(1.0), "b": [jnp.array(2.0), jnp.array(3.0)]}
jax.tree.map(lambda v: v * 2, tree)
Question 1 of 4