JAX 可微分编程框架
JAX 可微分编程框架:从零开始的完整教程
JAX 是一个结合了 NumPy 的便利性、自动微分和加速器(GPU/TPU)支持的数值计算库。它通过函数式编程和即时编译(JIT)实现高性能计算,特别适合机器学习、科学计算和可微分编程。
本教程将带你从基础概念到高级应用,全面掌握 JAX。
1. 为什么选择 JAX?
JAX 的核心优势在于:
- NumPy 兼容的 API:几乎可以无缝替换
numpy,且许多函数直接位于jax.numpy下。 - 自动微分:通过
grad、jacfwd、jacrev等函数变换,可对任意 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.cond、jax.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 可指定输出的布局。
结合 jit:jax.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 的强大之处在于可将 jit、grad、vmap 等自由组合。
示例:线性回归训练
# 生成数据
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.cond 和 jax.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 的可组合函数变换为科学计算和深度学习提供了一个灵活、高性能的框架。通过理解 jit、grad、vmap 和 pmap,你可以构建高效且清晰的可微分程序。