JAX 可微分编程框架

FreeGuideOnline 13阅读 2026-07-11

JAX 可微分编程框架:从零开始的完整教程

JAX 是一个结合了 NumPy 的便利性、自动微分和加速器(GPU/TPU)支持的数值计算库。它通过函数式编程和即时编译(JIT)实现高性能计算,特别适合机器学习、科学计算和可微分编程。

本教程将带你从基础概念到高级应用,全面掌握 JAX。

1. 为什么选择 JAX?

JAX 的核心优势在于:

  • NumPy 兼容的 API:几乎可以无缝替换 numpy,且许多函数直接位于 jax.numpy 下。
  • 自动微分:通过 gradjacfwdjacrev 等函数变换,可对任意 Python 函数求导。
  • XLA 即时编译 (JIT):使用 jit 将函数编译为高效的可执行代码,在 GPU/TPU 上大幅加速。
  • 自动向量化vmap 会自动为函数添加批次维度,避免手动编写循环。
  • Spmd 并行pmap 支持在多设备上并行计算。
  • 纯函数式风格:JAX 要求代码更接近纯函数,避免隐式副作用,这确保了变换的正确性。

JAX 与 NumPy 的关键区别:JAX 数组是不可变的,所有操作返回新数组;NumPy 数组则可原位修改。此外,JAX 使用 32 位浮点数作为默认值(在 GPU 上通常更快),而 NumPy 默认 64 位。

2. 安装与基础设置

pip install jax jaxlib
# 如需 GPU 支持,请根据 CUDA 版本安装对应 jaxlib,例如:
# pip install jax[cuda12]

导入常用模块:

import jax
import jax.numpy as jnp
from jax import grad, jit, vmap, pmap
import numpy as np  # 仍然可以导入原始 NumPy

启用 64 位精度(可选):

jax.config.update("jax_enable_x64", True)

3. JAX 数组:jax.numpy 基础

jax.numpy 创建和操作数组的方式几乎与 NumPy 完全相同。

# 创建数组
a = jnp.array([1.0, 2.0, 3.0])
b = jnp.zeros((2, 3))
c = jnp.arange(12).reshape(3, 4)

# 基本运算
print(a + 1)
print(jnp.dot(a, a))
print(jnp.sin(a))

# 随机数:JAX 需要显式管理随机种子
from jax import random
key = random.PRNGKey(42)
key, subkey = random.split(key)
x = random.normal(subkey, (3, 3))
print(x)

重要:JAX 数组的索引返回新数组,无法使用 x[idx] = value 在原处修改。应使用 x.at[idx].set(value) 来创建带更新的新数组:

x = jnp.array([1, 2, 3])
y = x.at[0].set(10)  # y = [10, 2, 3]  x 不变

4. 函数变换之一:jit — 即时编译

jit 将 Python 函数编译为 XLA 优化代码,通常在第一次调用时耗时较长(编译),后续调用极快。

def slow_function(x):
    for _ in range(1000):
        x = x + x * 0.0001
    return x

# 使用 jit 装饰器
@jit
def fast_function(x):
    for _ in range(1000):
        x = x + x * 0.0001
    return x

# 或直接调用 jit 转换
fast_function_alt = jit(slow_function)

# 测试速度
x = jnp.arange(1e6)
# slow_function(x)   # 较慢
# fast_function(x)   # 编译后极快

jit 的要求:函数内部的控制流不能依赖于输入数据的值。对于依赖数据的条件,可以使用 jax.lax.condjax.lax.while_loop 等。

5. 函数变换之二:grad — 自动微分

grad 返回一个计算函数标量输出梯度的新函数。

def f(x):
    return jnp.sum(x ** 2)

df = grad(f)      # f 关于 x 的梯度
x = jnp.array([1.0, 2.0, 3.0])
print(df(x))      # 输出 [2, 4, 6]

多参数函数:默认 grad 对第一个参数求导;使用 argnums 指定。

def loss(params, data):
    w, b = params
    return jnp.sum((jnp.dot(data, w) + b) ** 2)

# 对 params(第一个参数)求梯度
grad_loss = grad(loss, argnums=0)

params = (jnp.array([1.0, 2.0]), jnp.array(0.5))
data = jnp.array([[3.0, 4.0], [5.0, 6.0]])
print(grad_loss(params, data))

高阶导数:多次应用 grad

hessian = grad(grad(f))  # f 的 Hessian 向量积形式,或使用 jax.hessian

对非标量输出:使用 jax.jacfwd(前向模式)或 jax.jacrev(反向模式)计算雅可比矩阵。

def g(x):
    return jnp.array([x[0]**2, x[0]*x[1]])

j = jax.jacfwd(g)(jnp.array([2.0, 3.0]))
print(j)

6. 函数变换之三:vmap — 自动向量化

vmap 自动将函数映射到输入数组的批量维度上,避免编写显式循环,同时可结合 jit 获得最佳性能。

# 一个处理单一样本的函数
def single_predict(params, x):
    w, b = params
    return jnp.dot(w, x) + b

# 用 vmap 升级为批量版本
batch_predict = vmap(single_predict, in_axes=(None, 0))

params = (jnp.ones(3), 0.1)
xs = jnp.array([[1,2,3],[4,5,6],[7,8,9]])  # 形状 (3,3)
print(batch_predict(params, xs))            # 输出形状 (3,)

in_axes 指定每个输入参与映射的轴,None 表示不映射(广播)。out_axes 可指定输出的布局。

结合 jitjax.jit(vmap(f)) 或顺序相反,将使批量操作获得极致加速。

7. 函数变换之四:pmap — 多设备并行

pmap 在多个加速器(如多个 GPU)上并行计算,遵循 SPMD(单程序多数据)模型。

# 假设有 4 个设备,每个设备处理数据的一部分
def parallel_fn(x):
    return jnp.sum(x)

data = jnp.array([1.0, 2.0, 3.0, 4.0])
# 将数据分布到设备上需先使第一维等于设备数
data_parallel = data.reshape(4, 1)  # 形状 (4,1),4 对应设备数
pmapped_fn = pmap(parallel_fn)
result = pmapped_fn(data_parallel)  # 每个设备各自求和
print(result)  # 在 4 设备上输出 [1, 2, 3, 4]

一般需结合 jax.pmap 与数据并行训练策略使用。注意:pmap 下需要所有设备执行完全相同的指令。

8. 组合变换:构建可微分训练流水线

JAX 的强大之处在于可将 jitgradvmap 等自由组合。

示例:线性回归训练

# 生成数据
key = random.PRNGKey(0)
key, subkey = random.split(key)
X = random.normal(subkey, (100, 3))
true_w = jnp.array([2.0, -1.5, 4.0])
true_b = 0.3
y = jnp.dot(X, true_w) + true_b + 0.1 * random.normal(key, (100,))

# 模型和损失
def model(params, x):
    w, b = params
    return jnp.dot(x, w) + b

def loss_fn(params, x, y):
    preds = model(params, x)
    return jnp.mean((preds - y) ** 2)

# 梯度函数
grad_loss = jit(grad(loss_fn))

# 初始化参数
params = (jnp.zeros(3), 0.0)

# 训练循环
learning_rate = 0.1
for epoch in range(200):
    grads = grad_loss(params, X, y)
    # 更新参数(不可变更新)
    params = (params[0] - learning_rate * grads[0],
              params[1] - learning_rate * grads[1])
    if epoch % 50 == 0:
        loss = loss_fn(params, X, y)
        print(f"Epoch {epoch}, loss {loss:.4f}")

print("Estimated w:", params[0])
print("Estimated b:", params[1])

结合 vmap 应用:如果要对多个数据批次分别计算梯度,可以用 vmap(grad(loss_fn)) 实现 per-example gradients。

9. 状态管理与 flax 简介

JAX 函数是纯函数,不包含可变状态。对于有状态的操作(如批归一化统计、优化器动量),需要用显式状态传递或框架如 flax

手动管理参数状态:将参数传递给函数,返回更新后的参数,如上例所示。

常用深度学习库 flax 封装了模块定义和优化器状态,推荐用于复杂项目。

# Flax 简单示例(需安装 flax)
import flax.linen as nn
from flax.training import train_state
import optax

class MLP(nn.Module):
    @nn.compact
    def __call__(self, x):
        x = nn.Dense(32)(x)
        x = nn.relu(x)
        return nn.Dense(10)(x)

# 初始化
model = MLP()
key, subkey = random.split(random.PRNGKey(0))
x_dummy = jnp.ones((1, 784))
params = model.init(subkey, x_dummy)

# 优化器状态
tx = optax.adam(1e-3)
state = train_state.TrainState.create(apply_fn=model.apply, params=params, tx=tx)

# 训练步骤
def train_step(state, batch):
    def loss_fn(params):
        logits = model.apply(params, batch['image'])
        return optax.softmax_cross_entropy_with_integer_labels(logits, batch['label']).mean()
    grads = jax.grad(loss_fn)(state.params)
    state = state.apply_gradients(grads=grads)
    return state

10. 进阶技巧与常见陷阱

10.1 控制流:用 jax.lax 替代 Python 控制

JIT 编译要求函数内不能有数据依赖的 Python 条件/循环。使用 jax.lax.condjax.lax.fori_loop 等。

from jax import lax

def conditional_abs(x):
    return lax.cond(x >= 0, lambda t: t, lambda t: -t, operand=x)

10.2 动态形状与 jax.lax.slice

JIT 编译的数组形状必须在编译时已知(或通过静态参数传递)。需要动态索引时,使用 jax.lax.dynamic_slice 等。

10.3 避免在 jit 函数内调用 print

print 仅追踪时执行一次,非预期内。调试用 jax.debug.print

@jit
def debug_func(x):
    jax.debug.print("x = {x}", x=x)
    return x * 2

10.4 随机数正确用法

JAX 随机数生成器需要显式传递密钥。避免重用密钥,应使用 random.split

key = random.PRNGKey(0)
key, subkey = random.split(key)
x = random.normal(subkey, (5,))

10.5 性能分析

使用 with jax.profiler.trace(...) 和 TensorBoard 查看算子耗时。

jax.block_until_ready(result)  # 确保异步操作完成再计时

11. 完整项目模板

以下是一个典型的 JAX 研究项目结构(不含 Flax):

  • 定义模型:纯函数参数化。
  • 定义损失:标量输出。
  • 使用 jit(grad(loss)) 得到编译后的梯度计算。
  • 编写训练循环,通过 for 循环更新参数。
  • 可选:用 vmap 批量处理,用 pmap 多设备并行。

12. 资源与社区

JAX 的可组合函数变换为科学计算和深度学习提供了一个灵活、高性能的框架。通过理解 jitgradvmappmap,你可以构建高效且清晰的可微分程序。