JAX 入门:编译纯函数 本节摘要:PyTorch 即时修改变量,TensorFlow 构建计算图,JAX 编译纯函数——这最后一种彻底改变你对深度学习的思考方式。你已会用 PyTorch 搭网络:定义 、调 、走优化器。它能跑,千百万人在用。但 PyTorch 的 DNA 里烙着一个约束:它即时、逐个地在 Python 里跟踪运算,每次 都是一次独立内核启动,每步训练都重新解释同一份 Python 代码。这在你要跨 2048 块 TPU 训一个 5400 亿参数模型时会要你的命。
本节摘要:PyTorch 即时修改变量,TensorFlow 构建计算图,JAX 编译纯函数——这最后一种彻底改变你对深度学习的思考方式。你已会用 PyTorch 搭网络:定义
nn.Module、调.backward()、走优化器。它能跑,千百万人在用。但 PyTorch 的 DNA 里烙着一个约束:它即时、逐个地在 Python 里跟踪运算,每次tensor + tensor都是一次独立内核启动,每步训练都重新解释同一份 Python 代码。这在你要跨 2048 块 TPU 训一个 5400 亿参数模型时会要你的命。Google DeepMind 用 JAX 训 Gemini,Anthropic 用 JAX 训 Claude——这是地球上最大的神经网络训练任务,它们选 JAX,因为 JAX 把你的训练循环当作可编译的程序,而非一串 Python 调用。JAX 是带三大超能力的 NumPy:自动微分、JIT 编译到 XLA、自动向量化。你写一个处理单样本的函数,JAX 给你一个能处理整个 batch、算梯度、编译成机器码、跨多设备运行的函数——原函数一行都不用改。本节用 JAX + Optax 在 MNIST 上训一个 3 层 MLP,讲清 jax.numpy、jax.grad、jax.jit、jax.vmap 与函数式状态管理,并与 PyTorch 的面向对象方式对照。
阅读完本节,你应当能够:
你已会用 PyTorch 搭网络:定义 nn.Module、调 .backward()、走优化器。它能跑,千百万人在用。但 PyTorch 的 DNA 里烙着一个约束:它即时、逐个地在 Python 里跟踪运算,每次 tensor + tensor 都是一次独立内核启动,每步训练都重新解释同一份 Python 代码。这在你要跨 2048 块 TPU 训一个 5400 亿参数模型时会要你的命。
Google DeepMind 用 JAX 训 Gemini,Anthropic 用 JAX 训 Claude——这是地球上最大的神经网络训练任务。它们选 JAX,因为 JAX 把你的训练循环当作可编译的程序,而非一串 Python 调用。
JAX 是带三大超能力的 NumPy:自动微分、JIT 编译到 XLA、自动向量化。你写一个处理单样本的函数,JAX 给你一个能处理整个 batch、算梯度、编译成机器码、跨多设备运行的函数——原函数一行都不用改。
JAX 是函数式框架:没有类、没有可变状态、没有 .backward() 方法。
| PyTorch | JAX |
|---|---|
带 nn.Module 状态的类 |
纯函数:f(params, x) -> y |
loss.backward() |
jax.grad(loss_fn)(params, x, y) |
| 即时执行 | 经 XLA 的 JIT 编译 |
for x in batch: 手工循环 |
jax.vmap(f) 自动向量化 |
DataParallel / FSDP |
jax.pmap(f) 自动并行 |
可变的 model.parameters() |
不可变的数组 pytree |
这不是风格偏好,而是编译器约束。JIT 编译要求纯函数——相同输入永远产出相同输出,没有副作用。这个限制正是 100 倍提速成为可能的原因。
JAX 在加速器上重新实现了 NumPy API:函数名相同、广播规则相同、切片语义相同,但数组活在 GPU/TPU 上,每次运算都可被编译器跟踪。
import jax.numpy as jnp a = jnp.array([1.0, 2.0, 3.0]) b = jnp.array([4.0, 5.0, 6.0]) c = jnp.dot(a, b)
一个关键差别:JAX 数组不可变,没有 a[0] = 5,而是 a = a.at[0].set(5)。别扭一周就习惯了——不可变性正是让 grad、jit、vmap 可组合的根基。
PyTorch 把梯度挂在张量上(.grad),JAX 把梯度挂在函数上。
import jax def f(x): return x ** 2 df = jax.grad(f) df(3.0) # 6.0
jax.grad 接收一个函数,返回一个算梯度的新函数。没有 .backward() 调用,没有挂在张量上的计算图。梯度只是另一个你可以调用、组合、JIT 编译的函数。它可以任意组合:d2f = jax.grad(jax.grad(f))——二阶导、三阶导、雅可比、海森,全靠组合 grad。PyTorch 也能做(torch.autograd.functional.hessian),但那是后加的;在 JAX 里这是根基。约束:grad 只对纯函数工作——里面不能有 print(它们在跟踪时跑,而非执行时)、不能变异外部状态、不能用不带显式 key 的随机数。
@jax.jit def train_step(params, x, y): loss = loss_fn(params, x, y) return loss
首次调用时 JAX 跟踪函数——记录发生哪些运算但不执行,然后把跟踪结果交给 XLA(Accelerated Linear Algebra,Google 为 TPU 与 GPU 写的编译器)。XLA 融合运算、消除冗余内存拷贝、生成优化的机器码,后续调用完全绕开 Python,编译后的代码以 C++ 速度在加速器上跑。
JIT 何时有用:训练步(同样计算重复数千次)、推理(同模型不同输入)、任何被调用多次且输入形状相似的函数。JIT 何时有害:控制流依赖被跟踪数组值的函数(if x > 0 中 x 是被跟踪数组)、一次性计算(编译开销超运行时)、调试(跟踪隐藏真实执行)。控制流限制是真的:jax.lax.cond 取代 if/else,jax.lax.scan 取代 for 循环——这些不是可选项,是编译的代价。
你写一个处理单样本的函数:
def predict(params, x): return jnp.dot(params['w'], x) + params['b']
vmap 把它提升为处理整个 batch:
batch_predict = jax.vmap(predict, in_axes=(None, 0))
in_axes=(None, 0) 意思是:不对 params 批处理(共享),对 x 的轴 0 批处理。没有手工 for 循环、没有重塑、没有 batch 维穿线——JAX 自己找出 batch 维并把整个计算向量化。这不是语法糖,vmap 生成融合的向量化代码,比 Python 循环快 10~100 倍。而且它与 jit、grad 可组合:per_example_grads = jax.vmap(jax.grad(loss_fn), in_axes=(None, 0, 0))——逐样本梯度,一行,这在 PyTorch 里几乎不可能不靠 hack 实现。
parallel_step = jax.pmap(train_step, axis_name='devices')
pmap 把函数复制到所有可用设备(GPU/TPU)并切分 batch,函数内用 jax.lax.pmean、jax.lax.psum 跨设备同步梯度。Google 用 pmap(及其后继 shard_map)跨数千块 TPU v5e 芯片训 Gemini。编程模型:写单设备版,套 pmap,完事。
JAX 操作「pytree」——列表、元组、字典、数组的嵌套组合。你的模型参数就是一个 pytree。每个 JAX 变换(grad、jit、vmap)都知道如何遍历 pytree,jax.tree.map(f, tree) 把 f 作用到每个叶子上。这就是优化器一次更新全部参数的方式:
params = jax.tree.map(lambda p, g: p - lr * g, params, grads)
没有 .parameters() 方法、没有参数注册,树结构就是模型。
PyTorch 把状态存在对象里:
class Model(nn.Module): def __init__(self): self.linear = nn.Linear(784, 10) def forward(self, x): return self.linear(x)
JAX 用带显式状态的纯函数:
def predict(params, x): return jnp.dot(x, params['w']) + params['b']
params 是传进来的,什么都不存、什么都不变异。这让每个函数可测试、可组合、可编译,也意味着你要自己管 params——或用 Flax、Equinox 这样的库。
JAX 给原语,库给易用性:Flax(Google,神经网络层,带显式状态的 nn.Module)、Equinox(Patrick Kidger,基于 pytree、Pythonic)、Optax(DeepMind,优化器 + LR 调度,可组合的梯度变换)、Orbax(Google,检查点,存取 pytree)、CLU(Google,指标 + 日志)。Optax 是标准优化器库,它把梯度变换(Adam、SGD、裁剪)与参数更新解耦,让组合变得轻而易举:
optimizer = optax.chain( optax.clip_by_global_norm(1.0), optax.adam(learning_rate=1e-3), )
| 因素 | JAX | PyTorch |
|---|---|---|
| TPU 支持 | 一等公民(Google 造了两者) | 社区维护(torch_xla) |
| GPU 支持 | 良好(经 XLA 的 CUDA) | 业界最佳(原生 CUDA) |
| 调试 | 难(跟踪 + 编译) | 易(即时、逐行) |
| 生态 | 研究导向(Flax、Equinox) | 庞大(HuggingFace、torchvision 等) |
| 招聘 | 小众(Google/DeepMind/Anthropic) | 主流(到处都是) |
| 大规模训练 | 更优(XLA、pmap、mesh) | 良好(FSDP、DeepSpeed) |
| 原型速度 | 较慢(函数式开销) | 较快(变异即走) |
| 生产推理 | TensorFlow Serving、Vertex AI | TorchServe、Triton、ONNX |
| 谁在用 | DeepMind(Gemini)、Anthropic(Claude) | Meta(Llama)、OpenAI(GPT)、Stability AI |
诚实回答:除非有特定理由,否则用 PyTorch。那些理由是——有 TPU、需要逐样本梯度、超大规模多设备训练,或在 Google/DeepMind/Anthropic 工作。
JAX 没有全局随机状态,每次随机运算都要显式 PRNG key:
key = jax.random.PRNGKey(42) key1, key2 = jax.random.split(key) w = jax.random.normal(key1, shape=(784, 256))
一开始烦人,但它保证了跨设备与编译的可复现性——这是 PyTorch 的 torch.manual_seed 在多 GPU 下无法保证的。
我们用 JAX + Optax 在 MNIST 上训一个 3 层 MLP:784 输入、两个隐藏层(256、128)、10 个输出类。完整代码见原课程 phases/03-deep-learning-core/12-intro-to-jax/code/ 相应文件。
用 sklearn 取 MNIST,转成 jnp 数组。
没有类,只有一个返回 pytree 的函数。手工做 He 初始化,从一个种子 split 出三个 PRNG key,每个权重是一个不可变数组塞在嵌套字典里。
def init_params(key): k1, k2, k3 = random.split(key, 3) scale1 = jnp.sqrt(2.0 / 784) params = { 'layer1': {'w': scale1 * random.normal(k1, (784, 256)), 'b': jnp.zeros(256)}, # ... layer2 / layer3 同理 ... } return params
纯函数:params 进、预测出,没有 self、没有存储状态。loss_fn 从零算交叉熵——softmax、log、负均值。
def forward(params, x): x = jnp.dot(x, params['layer1']['w']) + params['layer1']['b'] x = jax.nn.relu(x) # ... layer2 / layer3 同理 ... return x def loss_fn(params, x, y): logits = forward(params, x) one_hot = jax.nn.one_hot(y, 10) return -jnp.mean(jnp.sum(jax.nn.log_softmax(logits) * one_hot, axis=-1))
@jax.jit def train_step(params, opt_state, x, y): loss, grads = jax.value_and_grad(loss_fn)(params, x, y) updates, opt_state = optimizer.update(grads, opt_state, params) params = optax.apply_updates(params, updates) return params, opt_state, loss
jax.value_and_grad 一次返回损失值与梯度。@jax.jit 把两个函数都编译到 XLA,首次调用后每步训练都不碰 Python。
optimizer = optax.adam(learning_rate=1e-3) key = random.PRNGKey(0) params = init_params(key) opt_state = optimizer.init(params) for epoch in range(n_epochs): key, subkey = random.split(key) perm = random.permutation(subkey, len(X_train)) # ... 取 batch、调 train_step ...
10 轮约 97% 测试精度。第 1 轮慢(JIT 编译),2~10 轮快。注意缺了什么:没有 .zero_grad()、没有 .backward()、没有 .step()——整个更新是一个组合的函数调用,梯度计算、Adam 变换、应用全在 train_step 里完成。
Flax 是最常见的 JAX 神经网络库,它把 nn.Module 加回来,但带显式状态管理:
import flax.linen as nn class MLP(nn.Module): @nn.compact def __call__(self, x): x = nn.Dense(256)(x); x = nn.relu(x) x = nn.Dense(128)(x); x = nn.relu(x) x = nn.Dense(10)(x) return x model = MLP() params = model.init(jax.random.PRNGKey(0), jnp.ones((1, 784))) logits = model.apply(params, x_batch)
结构与 PyTorch 相同,但 params 与模型分离:model.init() 创建 params,model.apply(params, x) 跑前向,模型对象本身无状态。
Equinox(Patrick Kidger)把模型表示为 pytree,模型本身就是 pytree,无需 .apply(),参数就是模型的叶子——更接近 JAX 的思考方式。
Optax 把梯度变换与更新解耦:梯度裁剪、学习率预热、权重衰减全都作为一串变换组合,每个变换看到梯度、修改它、传给下一个,没有臃肿的优化器类。
安装:
pip install jax jaxlib optax flax # GPU: pip install jax[cuda12] # TPU(Google Cloud): pip install jax[tpu] -f https://storage.googleapis.com/jax-releases/libtpu_releases.html
性能陷阱:首次 JIT 慢(编译),基准测试前要热身;JIT 内别用对 JAX 数组的 Python 循环,改用 jax.lax.scan 或 jax.lax.fori_loop;jax.debug.print() 在 JIT 内能工作,普通 print() 不行;JAX 默认预分配 75% GPU 显存,设 XLA_PYTHON_CLIENT_PREALLOCATE=false 可关。检查点用 Orbax。
本节产出(位于原课程 outputs/):
prompt-jax-optimizer.md:一个提示,帮你选对 JAX 优化器配置。skill-jax-patterns.md:一份覆盖 JAX 函数式模式的技能。jax.vmap 为 32 张 MNIST 图算逐样本梯度,算每个样本的梯度范数,哪些样本梯度最大?为什么?@jax.jit 的训练步,各计时 100 步——你的硬件上提速多大?首次调用的编译开销多少?a.at[0].set(5) 取代 a[0] = 5,别扭一周就习惯,不可变性让 grad/jit/vmap 可组合。.backward()、没有挂在张量上的计算图,梯度只是另一个可组合的函数,二阶/三阶导靠组合 grad。下一节(本章最后一节),我们讲调试神经网络——它编译了、它跑了、它产出了一个数,数是错的却什么都没崩,这是最难的调试:没有报错信息的那种。