jax-ml / jax
jax-ml/jax
可组合的Python+NumPy程序变换:微分、向量化、JIT到GPU/TPU等
项目概览
项目概述
JAX 是一个面向加速器的高性能数值计算与程序变换的 Python 库,由 Google 的 JAX 核心团队开发,以开源形式在 GitHub 上维护。它提供对 Python 和 NumPy 程序的可组合变换,包括自动微分、向量化和即时编译(JIT),并支持在 CPU、GPU 和 TPU 等硬件加速器上运行。JAX 的设计目标是让用户能够以接近 NumPy 的编程方式,获得高性能的数值计算和机器学习能力,同时通过可组合的变换系统支持复杂的科学计算和深度学习研究。
核心功能
- 自动微分(Automatic Differentiation):通过
jax.grad支持反向模式微分(反向传播),通过jax.jvp支持前向模式微分,且两者可以任意组合,支持高阶导数。 - 即时编译(JIT Compilation):通过
jax.jit将 Python 函数编译为 XLA 可执行程序,实现端到端的优化和加速,支持与自动微分、向量化等变换任意组合。 - 自动向量化(Automatic Vectorization):通过
jax.vmap自动将函数映射到数组的批次维度,将循环下沉到原始操作中,提升性能,避免手动重写批量代码。 - 并行计算(Parallel Computation):支持编译器自动并行化、显式分片(sharding)和手动设备级编程(
shard_map),可扩展到数千个设备。 - NumPy 兼容 API:提供
jax.numpy模块,覆盖大部分 NumPy 功能,并支持数组 API 标准。 - 可组合变换系统:JAX 的核心是可扩展的变换系统,允许用户自定义变换和自定义导数规则(
custom_jvp/custom_vjp)。 - 硬件加速支持:通过 XLA 编译器支持 CPU、NVIDIA GPU、Google TPU、AMD GPU 和 Apple GPU(实验性)等。
适用与不适用场景
适用场景:
- 需要高性能数值计算和科学计算的场景,如物理模拟、优化问题。
- 大规模机器学习模型的训练和推理,特别是需要 GPU/TPU 加速的场景。
- 需要自动微分和高阶导数的研究,如深度学习、最优化、概率编程。
- 需要将 NumPy 代码迁移到加速器上并提升性能的场景。
- 需要大规模并行计算(多设备、多主机)的场景。
不适用场景:
- 需要动态控制流(如依赖数据值的 Python 循环)且无法用 JAX 控制流原语(
lax.cond、lax.while_loop等)表达的场景,因为 JIT 编译要求静态形状。 - 需要与 TensorFlow 或 PyTorch 生态深度集成且无法通过 DLPack 等协议转换的场景。
- 对 Python 版本有严格限制且无法升级到 JAX 支持版本(当前最低 Python 3.11)的场景。
- 需要 Windows 上 NVIDIA GPU 支持(目前不支持)或 Apple GPU 稳定支持(目前实验性)的场景。
- 需要完全稳定的数值结果(JAX 不保证跨版本数值完全一致)的场景。
技术架构与依赖
JAX 的核心架构基于可组合的变换系统,底层依赖 XLA 编译器进行编译和优化,并通过 PJRT 运行时管理设备。主要组件包括:
- JAX 核心库:提供
jax模块,包含变换(grad、jit、vmap等)、数组类型(jax.Array)、NumPy 兼容 API(jax.numpy)等。 - XLA 编译器:用于将 JAX 程序编译为针对不同硬件的高效可执行代码,支持 StableHLO 作为输入。
- PJRT 运行时:提供统一的设备 API,支持插件式设备扩展。
- 依赖库:Python 3.11+、NumPy 2.0+、SciPy 1.14+、ml_dtypes 等。
安装与快速开始
安装:
根据平台选择安装命令(详见官方文档):
- CPU:
pip install -U jax - NVIDIA GPU:
pip install -U "jax[cuda13]" - Google TPU:
pip install -U "jax[tpu]" - AMD GPU (Linux):
pip install -U "jax[rocm7-local]" - Intel GPU:遵循 Intel 的指令
快速开始:
import jax
import jax.numpy as jnp
def predict(params, inputs):
for W, b in params:
outputs = jnp.dot(inputs, W) + b
inputs = jnp.tanh(outputs)
return outputs
def loss(params, inputs, targets):
preds = predict(params, inputs)
return jnp.sum((preds - targets)**2)
grad_loss = jax.jit(jax.grad(loss))
perex_grads = jax.jit(jax.vmap(grad_loss, in_axes=(None, 0, 0)))
典型使用方法
自动微分:
import jax
import jax.numpy as jnp
def tanh(x):
y = jnp.exp(-2.0 * x)
return (1.0 - y) / (1.0 + y)
grad_tanh = jax.grad(tanh)
print(grad_tanh(1.0)) # 0.4199743
JIT 编译:
import jax
import jax.numpy as jnp
def slow_f(x):
return x * x + x * 2.0
x = jnp.ones((5000, 5000))
fast_f = jax.jit(slow_f)
fast_f(x) # 编译并执行
自动向量化:
import jax
import jax.numpy as jnp
def l1_distance(x, y):
return jnp.sum(jnp.abs(x - y))
def pairwise_distances(dist1D, xs):
return jax.vmap(jax.vmap(dist1D, (0, None)), (None, 0))(xs, xs)
xs = jax.random.normal(jax.random.key(0), (100, 3))
dists = pairwise_distances(l1_distance, xs)
并行计算:
from jax.sharding import set_mesh, AxisType, PartitionSpec as P
mesh = jax.make_mesh((8,), ('data',), axis_types=(AxisType.Explicit,))
set_mesh(mesh)
# 参数分片、数据分片、自动并行化
配置与部署要点
- 64 位精度:默认使用 32 位浮点,可通过
jax.config.update('jax_enable_x64', True)启用 64 位。 - 设备选择:通过
JAX_PLATFORMS环境变量或jax.devices()控制。 - 内存管理:GPU 内存预分配可通过
XLA_PYTHON_CLIENT_MEM_FRACTION调整。 - 编译缓存:可配置持久化编译缓存以加速重复编译。
- 分布式部署:使用
jax.distributed.initialize()初始化多主机环境。 - 性能优化:使用
jax.jit包裹最外层函数,注意异步调度,使用.block_until_ready()同步。
限制、风险与许可证
限制:
- 不支持 Python 3.10 及以下版本(当前最低 3.11)。
- Windows 上不支持 NVIDIA GPU(仅实验性 WSL2)。
- Apple GPU 支持为实验性。
- 数值结果不保证跨版本完全一致。
- 动态形状支持有限,需要静态形状或使用形状多态。
风险:
- 项目处于积极开发中,API 可能变化,遵循 3 个月弃用政策。
- 部分功能标记为实验性,可能不稳定。
- 依赖 XLA 和 PJRT,这些外部项目的变更可能影响 JAX。
许可证: Apache-2.0
官方链接
- GitHub 仓库:https://github.com/jax-ml/jax
- 官方文档:https://docs.jax.dev/
- 变更日志:https://docs.jax.dev/en/latest/changelog.html
- PyPI 页面:https://pypi.org/project/jax/
信息来源和分析时间
本指南基于 GitHub 仓库 jax-ml/jax 的 README、CHANGELOG 和部分文档(如 docs/automatic-differentiation.md、docs/automatic-vectorization.md、docs/aot.md 等)编写。分析时间:2026-07-16(基于最新发布版本 JAX v0.11.0 的日期)。