Jax
A comprehensive skill catalog for AI agents
npx -y skills add G1Joshi/Agent-Skills --skill jaxAssembled from the repository path, not quoted from the project. Check it against their README if it does not work.
One thing to look at
- 10 stars10 stars. Stars are a popularity signal and not a quality one, but at this level it is likely that nobody has read this closely except its author, and you would be relying on your own review.
What its author says it does
Copied from the file, not written here
JAX high-performance numerical computing. Use for ML research.
SKILL.md
1.1 KB, as published. Nobody here has run it
JAX
JAX is "NumPy on steroids". It combines Autograd (automatic differentiation) with XLA (compilation). 2025 sees Flax NNX (PyTorch-style OOP) becoming standard.
When to Use
- TPU Training: JAX runs natively on Google TPUs.
- Research: If you need to compute 10th order derivatives or strange math.
- Massive Scale: DeepMind and OpenAI use JAX for training frontier models.
Core Concepts
Functional Transformations
grad(), jit(), vmap(), pmap().
Flax (NNX)
Neural network library. NNX introduces mutable state (OOP) to make JAX feel like PyTorch.
Statelessness
(Legacy Flax) parameters are stored separately from the model.
Best Practices (2025)
Do:
- Use
jit: Always compile your functions. - Use Flax NNX: Avoid the complexity of legacy immutable Flax/Haiku.
- Use
shard_map: For distributed training across devices.
Don't:
- Don't use side effects:
print()inside ajitfunction only runs once (during tracing).