简介
JAX 是一个面向加速器的高性能数组计算与程序变换的 Python 库,专为高性能数值计算和大规模机器学习而设计。它提供熟悉的 NumPy 风格 API,并内置可组合的函数变换,用于编译、批处理、自动微分和并行化。同一份代码可在 CPU、GPU 和 TPU 等多后端上运行,无需修改。
主要功能
- 熟悉的 NumPy 风格 API:降低学习成本,方便研究人员和工程师快速上手。
- 可组合的函数变换:包括
jit(即时编译)、vmap(自动向量化)、grad(自动微分)和 pmap(并行化),这些变换可以自由组合,简化复杂计算流程。
- 多后端运行:代码无需修改即可在 CPU、GPU 和 TPU 上执行,充分利用硬件加速能力。
- 高效的数组操作:专注于高性能数组运算,同时通过程序变换实现灵活优化。
生态系统
JAX 本身专注于核心的数组运算和程序变换,围绕它构建了一个不断发展的生态,涵盖机器学习与数值计算的多个领域:
- 神经网络:Flax、Equinox、Keras 等框架基于 JAX 构建。
- 优化器与求解器:Optax、Optimistix、Lineax、Diffrax 等提供优化与微分方程求解能力。
- 数据加载:Grain、TensorFlow Datasets、Hugging Face Datasets 等。
- 杂项工具:Orbax(存储)、Chex(测试工具)等。
- 概率编程:Blackjax、NumPyro、PyMC、TensorFlow Probability、Distrax 等。
- 物理与仿真:JAX MD、Brax 等用于分子动力学和机器人仿真。
- 大语言模型:MaxText、AXLearn、Levanter、EasyLM 等。
适用场景
- 高性能数值计算:科学计算、物理仿真、金融建模等需要大规模向量化运算的场景。
- 机器学习研究:快速原型设计、自定义训练循环、梯度计算与优化。
- 大规模深度学习:利用 GPU/TPU 加速训练,支持分布式并行。
- 自动微分与程序变换:需要将计算图编译优化、自动批处理或自动微分的复杂工作流。
快速开始
可通过官方文档的安装指南快速上手,并参考 JAX 101 教程 学习如何在 JAX 中思考,或查阅 JAX AI Stack 了解如何用 JAX 训练神经网络。