jax-ml/jax

Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more

解决的问题

JAX 提供了一个高性能的数值计算和大规模机器学习系统。开发者可以编写类似 NumPy 的代码,这些代码能够自动微分、编译以提升速度,并在 GPU 和 TPU 等多种硬件加速器上进行扩展。

工作原理

JAX 使用 XLA(加速线性代数)将 Python 和 NumPy 函数编译为优化的机器代码。它作为一个可组合的函数变换系统运行:

  • jax.grad:使用自动微分(支持反向模式和前向模式)计算 Python 和 NumPy 函数的梯度。
  • jax.jit:使用 XLA 将函数端到端编译,以实现更快的执行速度。
  • jax.vmap:自动向量化函数,将其映射到数组轴上,从而消除手动批处理循环。

适用人群

专为从事高性能数值计算和大规模机器学习的研究人员和开发者设计,适用于需要高效梯度计算以及在数千个设备上扩展计算能力的用户。

主要亮点

  • 可组合的变换:可任意顺序组合 gradjitvmap,创建高度优化的函数。
  • 硬件加速:通过 XLA 原生支持 NVIDIA GPU、Google TPU 及其他加速器。
  • 灵活的微分能力:可对循环、分支、递归和闭包进行任意阶数的微分。
  • 多种扩展方式:提供三种扩展模式:基于编译器的自动并行化、显式分片(sharding)和手动设备级编程。

相关

  • 项目
  • 项目
  • 项目
  • 项目
  • 项目