## 项目概述

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 的指令

**快速开始：**

```python
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)))
```

## 典型使用方法

**自动微分：**

```python
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 编译：**

```python
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)  # 编译并执行
```

**自动向量化：**

```python
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)
```

**并行计算：**

```python
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 的日期）。