Skip to content
← All tracks

JAX Performance & Distributed

JAX

jit details, memory and device control, sharding APIs (Mesh, PartitionSpec, shard_map), mixed precision. The performance and parallelism toolbox.

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

0 / 25 solved
  1. 1. Not solved yet. jit with static_argnums
  2. 2. Not solved yet. jit with static_argnames
  3. 3. Not solved yet. jit with donate_argnums
  4. 4. Not solved yet. jit with Pytree Input
  5. 5. Not solved yet. Explicit device_put
  6. 6. Not solved yet. block_until_ready for Sync
  7. 7. Not solved yet. Host Roundtrip via device_get
  8. 8. Not solved yet. make_jaxpr Inspection
  9. 9. Not solved yet. jit Retrace on Shape Change
  10. 10. Not solved yet. disable_jit Context Manager
  11. 11. Not solved yet. Mesh Creation
  12. 12. Not solved yet. PartitionSpec Basics
  13. 13. Not solved yet. NamedSharding (Replicated)
  14. 14. Not solved yet. shard_map Basics
  15. 15. Not solved yet. shard_map with psum Collective
  16. 16. Not solved yet. shard_map vs vmap
  17. 17. Not solved yet. bfloat16 Mixed Precision
  18. 18. Not solved yet. Explicit Dtype Promotion
  19. 19. Not solved yet. Fused Elementwise Ops
  20. 20. Not solved yet. with_sharding_constraint
  21. 21. Not solved yet. pjit Basics
  22. 22. Not solved yet. Compilation Cache Key
  23. 23. Not solved yet. compilation_cache.set_cache_dir
  24. 24. Not solved yet. Shape Polymorphism (Concept)
  25. 25. Not solved yet. Multi-Host JAX (Conceptual)

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

Where does a jnp array live by default?

jax.devices()
jax.default_backend()
jnp.ones(3).devices()
Question 1 of 4