We can't find the internet
Attempting to reconnect
Something went wrong!
Attempting to reconnect
← All tracks
0 / 4
JAX Foundations
JAXFunctional 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. Not solved yet. Vectorize with vmap
- 2. Not solved yet. Pure Function vs Impure
- 3. Not solved yet. Tracing: Shape vs Value
- 4. Not solved yet. JIT Compile a Function
- 5. Not solved yet. Jit with static_argnames
- 6. Not solved yet. Tracer Leak Detection
- 7. Not solved yet. Functional Update with .at[].set()
- 8. Not solved yet. Functional Update with .at[].add()
- 9. Not solved yet. 2-D Scatter via .at[].add()
- 10. Not solved yet. Gradient with jax.grad
- 11. Not solved yet. jax.value_and_grad
- 12. Not solved yet. jit + grad Composition
- 13. Not solved yet. PRNGKey and Split
- 14. Not solved yet. PRNGKey fold_in
- 15. Not solved yet. Deterministic Batch via vmap+split
- 16. Not solved yet. Pytree Leaves
- 17. Not solved yet. Pytree Map
- 18. Not solved yet. Gradient over Pytree Params
- 19. Not solved yet. JAX numpy vs PyTorch ops
- 20. Not solved yet. Dtype Promotion Rules
- 21. Not solved yet. take_along_axis (Gather)
- 22. Not solved yet. Multi-Condition where
- 23. Not solved yet. NaN-Safe Mean (with mask)
- 24. Not solved yet. Explicit Broadcasting via broadcast_to
- 25. Not solved yet. Clip and Extrema
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.
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