SKILL.md
JAX Skill
JAX is Autograd and XLA, brought together for high-performance machine learning research.
Contents
- [Concepts & Theory](reference.md)
- Immutability - The 4 Transformations - Pytrees
- [Code Examples](examples.md)
- jit, grad, vmap, random usage - Control Flow (scan, cond, fori_loop) - Parallelism (sharding)
Common Workflows
1. Developing a new Model
- Define your parameters as a Pytree (dict/dataclass).
- Define your forward pass function (pure).
- Define your loss function.
- Use
jax.valueandgradto get gradients. - Use
jax.jitto speed up the update step. - See [examples.md](examples.md) for snippets.
2. Debugging Shapes/NaNs
- Disable JIT:
jax.config.update("jaxdisablejit", True)to debug with standard python tools. - Use
jax.debug.printinside JITted functions.