JAX
常用谷歌的数值计算与深度学习框架,写法像 NumPy,能自动求导并编译加速。
JAX 是谷歌开源的 Python 数值计算库,接口接近 NumPy,另外提供几个可组合的函数变换:grad 自动求导、jit 用 XLA 编译器把代码编译成高效的 GPU/TPU 程序、vmap 自动批量化。Flax 是建在 JAX 上的神经网络库,负责定义网络层和管理参数,二者常一起出现。具身领域里,JAX 在大规模并行仿真和谷歌系工作中用得多:MJX、Brax 用它把物理仿真放到 GPU 上跑,Octo 和 Physical Intelligence 的 openpi 最初也是 JAX 实现。和 PyTorch 比,它更偏函数式写法,入门门槛略高。
例子openpi 仓库里的 π0 训练代码最早基于 JAX + Flax 编写。
- 也叫
- Flax
- 相关
- PyTorch、TensorFlow、MJX、Brax、openpi、Octo
- 来源
- jax-ml/jax GitHub
Flax 文档