JAX 入门:编译纯函数


文档摘要

JAX 入门:编译纯函数 本节摘要:PyTorch 即时修改变量,TensorFlow 构建计算图,JAX 编译纯函数——这最后一种彻底改变你对深度学习的思考方式。你已会用 PyTorch 搭网络:定义 、调 、走优化器。它能跑,千百万人在用。但 PyTorch 的 DNA 里烙着一个约束:它即时、逐个地在 Python 里跟踪运算,每次 都是一次独立内核启动,每步训练都重新解释同一份 Python 代码。这在你要跨 2048 块 TPU 训一个 5400 亿参数模型时会要你的命。

JAX 入门:编译纯函数

本节摘要: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 的面向对象方式对照。

学习目标

阅读完本节,你应当能够:

  1. 用 JAX 的函数式 API(jax.numpy、jax.grad、jax.jit、jax.vmap)编写纯函数式神经网络代码
  2. 解释 PyTorch 即时变异与 JAX 函数式编译模型之间的关键设计差异
  3. 应用 jit 编译与 vmap 向量化,加速训练循环对比朴素 Python。
  4. 在 JAX 里训练一个简单网络,并把显式状态管理与 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 的哲学

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:熟悉的外壳

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)。别扭一周就习惯了——不可变性正是让 gradjitvmap 可组合的根基。

jax.grad:函数式自动微分

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 的随机数。

jit:编译到 XLA

@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 循环——这些不是可选项,是编译的代价。

vmap:自动向量化

你写一个处理单样本的函数:

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 倍。而且它与 jitgrad 可组合:per_example_grads = jax.vmap(jax.grad(loss_fn), in_axes=(None, 0, 0))——逐样本梯度,一行,这在 PyTorch 里几乎不可能不靠 hack 实现。

pmap:跨设备数据并行

parallel_step = jax.pmap(train_step, axis_name='devices')

pmap 把函数复制到所有可用设备(GPU/TPU)并切分 batch,函数内用 jax.lax.pmeanjax.lax.psum 跨设备同步梯度。Google 用 pmap(及其后继 shard_map)跨数千块 TPU v5e 芯片训 Gemini。编程模型:写单设备版,套 pmap,完事。

Pytree:通用数据结构

JAX 操作「pytree」——列表、元组、字典、数组的嵌套组合。你的模型参数就是一个 pytree。每个 JAX 变换(gradjitvmap)都知道如何遍历 pytree,jax.tree.map(f, tree) 把 f 作用到每个叶子上。这就是优化器一次更新全部参数的方式:

params = jax.tree.map(lambda p, g: p - lr * g, params, grads)

没有 .parameters() 方法、没有参数注册,树结构就是模型。

函数式 vs 面向对象

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 生态

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 vs PyTorch

因素 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 里的随机数

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/ 相应文件。

Step 1:准备与数据

用 sklearn 取 MNIST,转成 jnp 数组。

Step 2:初始化参数

没有类,只有一个返回 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

Step 3:前向传播

纯函数: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))

Step 4:JIT 编译的训练步

@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。

Step 5:训练循环

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:Google 标准

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:Pythonic 替代

Equinox(Patrick Kidger)把模型表示为 pytree,模型本身就是 pytree,无需 .apply(),参数就是模型的叶子——更接近 JAX 的思考方式。

Optax:可组合优化器

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.scanjax.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 函数式模式的技能。

五、练习

  1. Easy:给 MLP 加 dropout——JAX 里 dropout 需要 PRNG key,把 key 穿过前向传播并为每个 dropout 层 split,比较有/无 dropout 的测试精度。
  2. Medium:用 jax.vmap 为 32 张 MNIST 图算逐样本梯度,算每个样本的梯度范数,哪些样本梯度最大?为什么?
  3. Hard:基准比较有/无 @jax.jit 的训练步,各计时 100 步——你的硬件上提速多大?首次调用的编译开销多少?

本节要点回顾

  1. JAX 编译纯函数,PyTorch 即时修改变量:这是根本设计差异,函数式约束(无副作用)是 100 倍提速的前提。
  2. JAX 数组不可变:a.at[0].set(5) 取代 a[0] = 5,别扭一周就习惯,不可变性让 grad/jit/vmap 可组合。
  3. jax.grad 把梯度挂在函数上:没有 .backward()、没有挂在张量上的计算图,梯度只是另一个可组合的函数,二阶/三阶导靠组合 grad。
  4. jit 跟踪 + 编译到 XLA:首次跟踪、XLA 融合运算生成机器码,后续绕开 Python;但要求纯函数,值依赖的控制流要用 lax.cond / lax.scan。
  5. vmap 自动向量化:写单样本函数,vmap 提升为 batch 版,生成融合代码快 10~100 倍,与 grad/jit 可组合,逐样本梯度一行搞定。
  6. pmap 跨设备自动并行:写单设备版、套 pmap 即可,DeepMind 训 Gemini 就这么跨数千 TPU。
  7. pytree 是通用数据结构:模型参数是嵌套字典/列表/数组,每个变换都能遍历,树结构就是模型,无需 .parameters()。
  8. Optax 把梯度变换解耦:裁剪/Adam/调度组合成链,每个变换看梯度、改、传,无臃肿优化器类。
  9. 随机数需显式 key:无全局随机状态,保证跨设备/编译可复现。
  10. 默认用 PyTorch:除非有 TPU、需逐样本梯度、超大规模多设备或在 Google/DeepMind/Anthropic,否则 PyTorch 更主流。

下一节(本章最后一节),我们讲调试神经网络——它编译了、它跑了、它产出了一个数,数是错的却什么都没崩,这是最难的调试:没有报错信息的那种。


发布者: 作者: Rohit Gupta 转发
评论区 (0)
U