JAX
CommonGoogle's NumPy-like framework for numerical computing and deep learning, with automatic differentiation and compilation.
JAX is Google's open-source Python library for numerical computing, with an interface close to NumPy plus a set of composable function transforms: grad for automatic differentiation, jit for compiling code into fast GPU/TPU programs via the XLA compiler, and vmap for automatic batching. Flax is a neural-network library built on top of JAX that defines layers and manages parameters; the two are usually used together. In embodied AI, JAX shows up heavily in large-scale parallel simulation and in Google-adjacent work: MJX and Brax use it to run physics simulation on GPUs, and Octo and Physical Intelligence's openpi were originally implemented in JAX. Compared with PyTorch, it leans more functional in style and has a somewhat steeper learning curve.
ExampleThe π0 training code in the openpi repository was originally written in JAX and Flax.
- Also called
- Flax
- Related
- PyTorch · TensorFlow · MuJoCo XLA · Brax · openpi (Physical Intelligence) · Octo
- Sources
- jax-ml/jax GitHub
Flax 文档 (Chinese)