We can't find the internet
Attempting to reconnect
Something went wrong!
Attempting to reconnect
← All tracks
0 / 4
JAX Conceptual Deep-Dives
JAXCustom 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. Not solved yet. register_pytree_node
- 2. Not solved yet. tree_flatten / tree_unflatten Round-Trip
- 3. Not solved yet. Pytree Aux Data
- 4. Not solved yet. pure_callback Basics
- 5. Not solved yet. io_callback for Side Effects
- 6. Not solved yet. jax.debug.print
- 7. Not solved yet. JAX <-> NumPy Bridge
- 8. Not solved yet. int8 Affine Quantization
- 9. Not solved yet. int8 Dequantization
- 10. Not solved yet. float8 Round-Trip
- 11. Not solved yet. Counting Jaxpr Equations
- 12. Not solved yet. jax.debug.callback
- 13. Not solved yet. Selective stop_gradient (Frozen Layers)
- 14. Not solved yet. jvp with Structured Tangents
- 15. Not solved yet. jaxpr with jit
Check yourself
4 questions · one attempt eachThese do not count toward finishing the track. They are here to catch the things that are easy to read past.
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