jax-ml / jax

jax-ml/jax

open_in_new前往仓库

可组合的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 仓库 jax-ml/jax 的 README、CHANGELOG 和部分文档(如 docs/automatic-differentiation.md、docs/automatic-vectorization.md、docs/aot.md 等)编写。分析时间:2026-07-16(基于最新发布版本 JAX v0.11.0 的日期)。