AI 工程基础体系 · 第 69/100 篇。内容覆盖机器学习、深度学习与生成式 AI;模型、数据、评测、权限和成本会作为同一生产系统处理。

JAX 工程基础:JIT、Grad、Vmap、Pmap、纯函数和设备执行

JAX 是一个以数组计算为核心的 Python 框架。它使用接近 NumPy 的数组接口,并通过程序变换(transformation)把同一个函数转换成:

  • 可编译执行的函数:jax.jit
  • 可求导的函数:jax.grad
  • 自动批处理的函数:jax.vmap
  • 跨设备并行的函数:jax.pmap

这些能力并不是四套彼此独立的 API。它们共同依赖几个基础前提:

  1. 计算最好表示为纯函数;
  2. 函数的输入、输出和中间状态应当显式传递;
  3. 数组的形状、数据类型和设备位置会影响编译与执行;
  4. 变换通常先追踪 Python 函数,再生成或组合新的计算程序;
  5. JAX 数组的计算可能异步提交到设备,提交完成不等于计算已经完成。

理解这些前提,比记住某个装饰器的写法更重要。


一、先建立一个 JAX 的执行模型

一个普通 Python 函数通常在调用时立即执行:

def add_one(x):
    return x + 1

y = add_one(3)

JAX 变换后的函数并不一定直接按照 Python 解释器的逐行语义执行。以 jax.jit 为例,第一次调用时大致会经历:

Python 函数
   │
   ├─ 使用抽象值进行追踪(tracing)
   │
   ├─ 生成 JAX 运算表达式
   │
   ├─ 编译为后端可执行程序
   │
   └─ 在 CPU、GPU 或 TPU 上执行

后续调用如果满足相同的编译条件,通常可以复用已有编译结果。

这里的“相同编译条件”至少与以下因素有关:

  • 数组的形状;
  • 数据类型;
  • 静态参数的值;
  • 被变换函数的结构;
  • 某些控制流和配置。

因此,JAX 不是“把任意 Python 程序自动变快”。它更适合将计算表达为形状相对稳定、数据通过参数传入、控制流可被 JAX 表达的数值函数。


二、纯函数:JAX 变换能够成立的基础

2.1 什么是纯函数

如果一个函数满足:

y=f(x)y = f(x)

那么在相同输入 xx 下,它总是产生相同输出 yy,并且不会依赖或修改函数外部的隐藏状态,这个函数就是纯函数。

例如:

import jax.numpy as jnp

def linear(params, x):
    return x @ params["w"] + params["b"]

这个函数的结果只取决于 paramsx

下面的函数则依赖隐藏状态:

running_count = 0

def bad_function(x):
    global running_count
    running_count += 1
    return x + running_count

在普通 Python 中,它可能返回:

第一次调用:x + 1
第二次调用:x + 2

但如果把它交给 jax.jit,JAX 追踪时执行的 Python 副作用和实际设备执行并不是同一个时刻。计数器可能只在追踪或重新编译时修改,而不是每次设备执行时修改。

因此,不应在被 jitgradvmappmap 变换的函数中依赖:

  • 全局可变变量;
  • print 的执行次数;
  • 文件写入;
  • 随意修改 Python 容器;
  • 隐式的随机数状态;
  • 隐式的训练模式或评估模式状态。

副作用不是“绝对禁止”,而是不能把它当作 JAX 数值计算语义的一部分。

2.2 JAX 数组的不可变语义

JAX 数组采用不可变语义。下面的代码不是原地修改:

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

print(x)  # [1 2 3]
print(y)  # [ 1 10  3]

x.at[1].set(10) 表示产生一个新数组。JAX 编译器在确定安全时,可能在底层复用内存,但这种优化不会改变不可变语义。

这与纯函数的关系是:函数通过返回新值表达状态变化,而不是修改外部对象。

例如优化器状态可以显式写成:

def sgd_update(params, grads, learning_rate):
    return jax.tree_util.tree_map(
        lambda p, g: p - learning_rate * g,
        params,
        grads,
    )

params 没有被原地修改,更新后的参数作为返回值传给下一步。

2.3 PyTree:显式携带复杂状态

JAX 的 PyTree 是由字典、元组、列表等容器组成的树,其叶子通常是数组。

下面的参数就是一个 PyTree:

params = {
    "w": jnp.ones((4, 2)),
    "b": jnp.zeros((2,)),
}

jax.tree_util.tree_map 会对结构中对应位置的叶子执行相同操作:

new_params = jax.tree_util.tree_map(
    lambda p, g: p - 0.01 * g,
    params,
    grads,
)

这使得模型参数、梯度、优化器状态、统计量都可以作为函数输入和输出显式流动:

(params, optimizer_state, batch, rng_key)
        │
        ▼
(params_new, optimizer_state_new, metrics)

这比把状态放在对象属性或全局变量中更适合 JAX 变换,也更便于检查点保存、恢复、评测和多设备复制。

2.4 随机数也必须显式传递

JAX 不使用隐式全局随机状态。随机数函数接收一个 key,并返回随机结果:

import jax

key = jax.random.key(0)
key, subkey = jax.random.split(key)
x = jax.random.normal(subkey, (3,))

一个 key 不应被无限重复使用。更清晰的写法是每次使用前拆分:

def sample_step(key, shape):
    key, subkey = jax.random.split(key)
    sample = jax.random.normal(subkey, shape)
    return key, sample

这里的 keysample 都是显式输出。随机性仍然存在,但从函数接口看,它表现为:

(key,input)(new_key,random_output)(\text{key}, \text{input}) \rightarrow (\text{new\_key}, \text{random\_output})

这样做的好处包括:

  • jit 可以正确追踪随机计算;
  • vmap 可以为每个样本生成不同 key;
  • pmap 可以为每个设备生成独立 key;
  • 训练可以精确恢复或复现。

三、jax.jit:把纯数值函数编译到设备

3.1 jit 做了什么

jax.jit(f) 返回一个经过即时编译的函数:

import jax
import jax.numpy as jnp

def square_sum(x):
    return jnp.sum(x * x)

compiled_square_sum = jax.jit(square_sum)

x = jnp.arange(5.0)
y = compiled_square_sum(x)
print(y)  # 30.0

第一次调用通常包含追踪和编译开销,后续相同形状和类型的调用可以复用编译结果。

jit 的价值不是单纯减少 Python 函数调用次数,而是把一组数组运算融合成更适合设备执行的程序。例如:

def f(x, w, b):
    z = x @ w
    z = z + b
    return jnp.tanh(z)

如果每一步都在 Python 层单独调度,可能产生多个设备操作;经过 jit 后,编译器可以将其作为一个整体优化。

3.2 完整的 JIT 训练步

下面实现一个线性回归训练步:

y^=XW+b\hat{y} = XW + b

均方误差为:

L(W,b)=1Ni=1N(y^iyi)2L(W,b) = \frac{1}{N}\sum_{i=1}^{N} (\hat{y}_i-y_i)^2

代码如下:

import jax
import jax.numpy as jnp

def predict(params, x):
    return x @ params["w"] + params["b"]

def mse_loss(params, x, y):
    pred = predict(params, x)
    return jnp.mean((pred - y) ** 2)

def train_step(params, x, y, learning_rate):
    loss, grads = jax.value_and_grad(mse_loss)(params, x, y)

    new_params = jax.tree_util.tree_map(
        lambda p, g: p - learning_rate * g,
        params,
        grads,
    )
    return new_params, loss

jit_train_step = jax.jit(train_step)

key = jax.random.key(0)
key, kx, kw = jax.random.split(key, 3)

x = jax.random.normal(kx, (32, 4))
true_w = jnp.array([
    [1.0, -2.0],
    [0.5,  1.0],
    [2.0,  0.0],
    [-1.0, 0.5],
])
true_b = jnp.array([0.3, -0.7])
y = x @ true_w + true_b

params = {
    "w": jax.random.normal(kw, (4, 2)),
    "b": jnp.zeros((2,)),
}

for step in range(5):
    params, loss = jit_train_step(params, x, y, 0.05)
    print(step, loss)

输入包括:

  • params:模型参数 PyTree;
  • x:形状为 (32, 4) 的批量输入;
  • y:形状为 (32, 2) 的目标;
  • learning_rate:标量学习率。

输出包括:

  • 更新后的参数;
  • 当前批次损失。

这里 value_and_grad 同时计算函数值和梯度,避免为同一个损失重复组织计算。jit 包住整个训练步,使前向、反向和参数更新都成为一个可编译的计算单元。

3.3 静态参数与动态参数

设备计算通常要求数组形状在编译时可知。如果函数中的参数只影响数值,可以作为动态数组传入:

@jax.jit
def scale(x, factor):
    return x * factor

scale(jnp.ones(4), 2.0)

但如果参数改变了 Python 控制流或输出结构,通常需要标记为静态参数:

from functools import partial

@partial(jax.jit, static_argnames=("normalize",))
def transform(x, normalize):
    if normalize:
        return x / jnp.linalg.norm(x)
    return x

normalize 是 Python 布尔值,决定了要追踪哪一条 Python 分支,所以将其标记为静态参数。

静态参数的每个不同值都可能导致一次新的编译:

transform(x, True)   # 一种编译版本
transform(x, False)  # 另一种编译版本

如果把不断变化的大对象标记为静态参数,会造成:

  • 编译缓存快速膨胀;
  • 编译时间增加;
  • 哈希或比较静态参数失败;
  • 代码在长时间运行后性能恶化。

经验上,静态参数适合表示少量配置,例如层数、是否归一化、算法分支;不适合表示每个 batch 都变化的数组数据。

3.4 JIT 下常见的追踪错误

jit 追踪时,函数参数可能是 tracer,而不是普通 Python 数值。下面的代码可能失败:

@jax.jit
def bad(x):
    if x > 0:
        return x
    return -x

原因是 x > 0 产生的结果也是数组计算结果,不能直接作为 Python if 的布尔条件。可以改为:

@jax.jit
def good(x):
    return jnp.where(x > 0, x, -x)

或者使用 JAX 控制流:

@jax.jit
def good_cond(x):
    return jax.lax.cond(
        x > 0,
        lambda value: value,
        lambda value: -value,
        x,
    )

如果分支条件是静态 Python 配置,则可以使用前面的 static_argnames

类似地,以下操作通常会破坏追踪:

@jax.jit
def bad_to_numpy(x):
    return float(x)

设备数组转换为 Python 标量需要把值取回主机,而追踪阶段通常没有实际数值可供转换。应将数值转换放在 JIT 函数外部。


四、jax.grad:对标量函数进行自动微分

4.1 梯度的数学含义

给定标量函数:

L(θ)L(\theta)

其中 θ\theta 是参数,梯度定义为:

θL=[Lθ1,Lθ2,]\nabla_\theta L = \left[ \frac{\partial L}{\partial \theta_1}, \frac{\partial L}{\partial \theta_2}, \ldots \right]

梯度指向函数增长最快的方向。因此梯度下降使用:

θnew=θηθL\theta_{\text{new}} = \theta - \eta \nabla_\theta L

其中:

  • η\eta 是学习率;
  • θL\nabla_\theta L 是损失对参数的梯度。

jax.grad 对可由 JAX 运算构成的函数执行自动微分:

import jax
import jax.numpy as jnp

def loss_fn(w, x, y):
    prediction = w * x
    return jnp.mean((prediction - y) ** 2)

x = jnp.array([1.0, 2.0, 3.0])
y = jnp.array([2.0, 4.0, 6.0])
w = 1.0

grad_fn = jax.grad(loss_fn)
print(grad_fn(w, x, y))

loss_fn 必须返回标量,或者至少对求导的那一维产生可定义的标量目标。直接对向量输出使用 jax.grad 通常会报错:

def vector_output(w, x):
    return w * x

# jax.grad(vector_output)(w, x)  # 不适用于向量输出

如果需要完整雅可比矩阵,可以使用 jax.jacfwdjax.jacrev

jacobian_fn = jax.jacrev(vector_output)
jacobian = jacobian_fn(w, x)

4.2 从标量损失到参数 PyTree

grad 不要求参数是单个数组。它可以对 PyTree 参数求导:

grads = jax.grad(mse_loss)(params, x, y)
print(grads["w"].shape)  # (4, 2)
print(grads["b"].shape)  # (2,)

梯度树与参数树结构一致:

params:
  w: (4, 2)
  b: (2,)

grads:
  w: (4, 2)
  b: (2,)

这就是前面可以使用 tree_map 更新参数的原因。

4.3 value_and_grad 和辅助输出

训练时通常既需要损失,又需要梯度:

loss, grads = jax.value_and_grad(mse_loss)(params, x, y)

如果损失函数还要返回预测值、准确率等辅助信息,可以使用 has_aux=True

def loss_with_metrics(params, x, y):
    pred = predict(params, x)
    loss = jnp.mean((pred - y) ** 2)
    mae = jnp.mean(jnp.abs(pred - y))
    return loss, {"mae": mae}

(loss, metrics), grads = jax.value_and_grad(
    loss_with_metrics,
    has_aux=True,
)(params, x, y)

这里要求返回值结构是:

(scalar_loss, auxiliary_output)

梯度只针对第一个标量损失计算,辅助输出不会被当作求导目标。

4.4 反向模式与前向模式的取舍

jax.grad 通常基于反向模式自动微分。对神经网络而言,参数数量往往远大于输出标量,因此反向模式适合计算:

Lθ\frac{\partial L}{\partial \theta}

但如果输入维度很小、输出维度很多,前向模式可能更合适。JAX 提供:

  • jax.jacfwd:前向模式构造雅可比;
  • jax.jacrev:反向模式构造雅可比;
  • jax.jvp:计算方向导数;
  • jax.vjp:获取向量-雅可比积。

grad 不是“任何函数都能求导”。Python 字符串处理、任意文件操作、未经 JAX 注册的外部函数,都不自动具有可微分规则。


五、jax.vmap:把单样本函数提升为批量函数

5.1 从单样本函数开始

设单个样本为 xix_i,模型产生:

f(xi)f(x_i)

如果有 NN 个样本:

X=[x1,x2,,xN]X = [x_1, x_2, \ldots, x_N]

批量计算希望得到:

[f(x1),f(x2),,f(xN)][f(x_1), f(x_2), \ldots, f(x_N)]

vmap 自动完成这种向量化,而不需要手写 Python for 循环。

import jax
import jax.numpy as jnp

def single_prediction(params, x):
    return x @ params["w"] + params["b"]

batch_prediction = jax.vmap(
    single_prediction,
    in_axes=(None, 0),
    out_axes=0,
)

这里:

  • params 不沿 batch 维映射,因此是 None
  • x 的第 0 维是 batch 维,因此是 0
  • 输出沿第 0 维堆叠,因此 out_axes=0

调用:

params = {
    "w": jnp.ones((4, 2)),
    "b": jnp.zeros((2,)),
}
x_batch = jnp.ones((8, 4))

pred = batch_prediction(params, x_batch)
print(pred.shape)  # (8, 2)

5.2 vmap 与手写循环的差异

手写循环版本是:

def loop_prediction(params, x_batch):
    return jnp.stack([
        single_prediction(params, x)
        for x in x_batch
    ])

它在 Python 层展开循环。vmap 则把“沿某个轴批量应用函数”作为变换交给 JAX:

batch_prediction = jax.vmap(
    single_prediction,
    in_axes=(None, 0),
)

这不仅让代码更短,更重要的是它保留了计算结构,使其可以继续与 jitgrad 等组合:

compiled_batch_prediction = jax.jit(batch_prediction)

5.3 用 vmap 计算逐样本梯度

设单个样本损失为:

i(θ)\ell_i(\theta)

整个批次的损失为:

L(θ)=1Ni=1Ni(θ)L(\theta) = \frac{1}{N}\sum_{i=1}^{N}\ell_i(\theta)

可以先写单样本损失:

def single_loss(params, x, y):
    pred = single_prediction(params, x)
    return jnp.mean((pred - y) ** 2)

对单样本求梯度,再沿样本维度映射:

per_example_grad = jax.vmap(
    jax.grad(single_loss),
    in_axes=(None, 0, 0),
)

grads_each_sample = per_example_grad(
    params,
    x_batch,
    y_batch,
)

结果中每个参数叶子都会增加一个样本维度。例如:

grads_each_sample["w"].shape == (batch_size, 4, 2)
grads_each_sample["b"].shape == (batch_size, 2)

再对样本维度求平均,就得到批量梯度:

mean_grads = jax.tree_util.tree_map(
    lambda g: jnp.mean(g, axis=0),
    grads_each_sample,
)

这与直接对批量平均损失求梯度在数学上应当一致:

def batch_loss(params, x_batch, y_batch):
    losses = jax.vmap(single_loss, in_axes=(None, 0, 0))(
        params,
        x_batch,
        y_batch,
    )
    return jnp.mean(losses)

direct_grads = jax.grad(batch_loss)(
    params,
    x_batch,
    y_batch,
)

逐样本梯度常用于:

  • 差分隐私训练;
  • 样本级梯度裁剪;
  • 影响函数分析;
  • 样本筛选和诊断。

但它会保留 batch 维度,可能显著增加内存占用。对于普通训练,直接计算批量损失的梯度通常更节省内存。

5.4 vmap 不等于多设备并行

vmap 的主要语义是:

一个设备上的向量化批量计算

它通常将标量或单样本计算转换为数组批量计算,并不承诺将不同样本分配到不同设备。

因此:

  • 想在单个设备上处理 batch,优先考虑 vmap 或直接使用批量数组;
  • 想把计算映射到多个设备,并使用设备间通信,考虑 pmap
  • 想表达更一般的数组分片和设备布局,需要进一步理解 JAX 的 sharding 体系。

六、pmap:跨设备映射与集体通信

6.1 pmap 的语义

pmap 将一个函数映射到多个设备上。假设有 DD 个设备,输入的第 0 维大小为 DD

x.shape = (D, local_batch, ...)

那么第 dd 个设备接收:

x[d]

函数在每个设备上执行一次。与 vmap 不同,pmap 还可以使用设备间集体通信,例如:

  • lax.pmean:跨设备求平均;
  • lax.psum:跨设备求和;
  • lax.pmax:跨设备求最大值;
  • lax.all_gather:收集各设备数据。

pmap 中的通信必须通过 axis_name 标识的映射轴完成:

jax.pmap(step, axis_name="devices")

函数内部才能使用:

jax.lax.pmean(value, axis_name="devices")

6.2 数据并行训练的完整例子

下面实现一个可运行的数据并行训练步。每个设备处理一个局部 batch,先计算局部梯度,再通过 pmean 求全局平均梯度。

import jax
import jax.numpy as jnp

def predict(params, x):
    return x @ params["w"] + params["b"]

def mse_loss(params, x, y):
    pred = predict(params, x)
    return jnp.mean((pred - y) ** 2)

def local_train_step(params, x, y, learning_rate):
    loss, grads = jax.value_and_grad(mse_loss)(params, x, y)

    # 在所有参与 pmap 的设备上求梯度平均
    grads = jax.lax.pmean(grads, axis_name="devices")

    # 每个设备使用相同的全局平均梯度更新自己的参数副本
    new_params = jax.tree_util.tree_map(
        lambda p, g: p - learning_rate * g,
        params,
        grads,
    )

    # 记录全局平均损失,而不是只记录本设备损失
    mean_loss = jax.lax.pmean(loss, axis_name="devices")
    return new_params, mean_loss

parallel_train_step = jax.pmap(
    local_train_step,
    axis_name="devices",
)

准备数据时,需要使第 0 维等于设备数:

n_devices = jax.local_device_count()
print("local devices:", n_devices)

n_samples = n_devices * 8
n_features = 4
n_outputs = 2

key = jax.random.key(42)
key, kx, kw = jax.random.split(key, 3)

x = jax.random.normal(kx, (n_samples, n_features))

true_w = jnp.array([
    [1.0, -2.0],
    [0.5,  1.0],
    [2.0,  0.0],
    [-1.0, 0.5],
])
true_b = jnp.array([0.3, -0.7])
y = x @ true_w + true_b

params = {
    "w": jax.random.normal(kw, (n_features, n_outputs)),
    "b": jnp.zeros((n_outputs,)),
}

pmap 需要每个设备都有一份参数,因此将参数复制到所有本地设备:

replicated_params = jax.device_put_replicated(
    params,
    jax.local_devices(),
)

将 batch 按设备切分:

per_device_batch = n_samples // n_devices

x_sharded = x.reshape(
    n_devices,
    per_device_batch,
    n_features,
)
y_sharded = y.reshape(
    n_devices,
    per_device_batch,
    n_outputs,
)

执行训练:

for step in range(5):
    replicated_params, losses = parallel_train_step(
        replicated_params,
        x_sharded,
        y_sharded,
        0.05,
    )

    # losses 的形状为 (n_devices,)
    # 由于已经 pmean,每个设备上的值应当相同
    print(step, float(losses[0]))

数据流可以表示为:

全局 batch
    │
    ├── 设备 0:x[0], y[0] ──┐
    ├── 设备 1:x[1], y[1] ──┤
    ├── ...                  ├─ 局部前向与局部反向
    └── 设备 D-1             │
                             ▼
                 lax.pmean(grads, "devices")
                             │
                             ▼
                 每个设备得到相同全局梯度
                             │
                             ▼
                 各设备独立更新自己的参数副本

如果没有 pmean,每个设备会使用自己的局部梯度更新参数,模型副本会逐步产生差异。这不是标准的数据并行同步训练。

6.3 为什么要同时对梯度和损失做 pmean

假设每个设备的局部 batch 大小相同,设备 dd 的局部平均损失为:

Ld=1BidiL_d = \frac{1}{B}\sum_{i \in d}\ell_i

全局平均损失为:

L=1Dd=1DLdL = \frac{1}{D}\sum_{d=1}^{D}L_d

所以可以使用:

jax.lax.pmean(loss, "devices")

同理,标准数据并行训练需要:

g=1Dd=1Dgdg = \frac{1}{D}\sum_{d=1}^{D}g_d

因此要对局部梯度执行 pmean

如果各设备的 batch 大小不相同,简单的 pmean 可能不再等于按样本数加权的全局平均。此时应传递有效样本数,并使用 psum 计算加权结果。例如:

g=dBdgddBdg = \frac{\sum_d B_d g_d}{\sum_d B_d}

这说明“每个设备做一个平均,再对设备平均”隐含了各设备样本数相等的条件。

6.4 pmap 的设备范围

jax.local_device_count() 返回当前进程可见的本地设备数量。常见情况是:

  • CPU 环境通常只有一个本地设备;
  • GPU 环境取决于可见 GPU 数量;
  • TPU 和多主机环境可能有本地设备与全局设备的区别。

pmap 默认主要面向当前进程可见的设备。多主机训练还需要:

  • 所有主机以一致的程序启动;
  • 正确初始化 JAX 的分布式环境;
  • 保证各主机输入切分和步骤同步;
  • 正确处理主机间故障、重启和检查点。

不能仅仅把单机的 pmap 代码复制到多主机环境,就假设它自动完成了完整的全局数据并行。

此外,JAX 的设备并行 API 在不同版本中持续演进。新的项目还应了解基于显式 sharding 的 API,例如 pjitshard_map 和相关 jax.sharding 能力。pmap 仍然适合说明和实现经典的复制式数据并行,但不应把它理解为 JAX 唯一或长期唯一的多设备抽象。


七、gradvmapjitpmap 如何组合

这些变换可以嵌套,但嵌套顺序决定语义。

7.1 jit(grad(f))

compiled_grad = jax.jit(jax.grad(loss_fn))

含义是:

  1. 先构造损失函数的梯度函数;
  2. 再将梯度计算编译到设备。

这是训练中非常常见的组合。

7.2 grad(jit(f))

grad_of_compiled = jax.grad(jax.jit(loss_fn))

在很多纯 JAX 数值函数上也可以工作,因为 jit 本身可以被自动微分规则处理。但工程上通常更容易理解和控制的是:

jax.jit(jax.grad(loss_fn))

也就是先明确数学函数和梯度,再编译最终训练步。

7.3 vmap(grad(f))

per_example_grad = jax.vmap(
    jax.grad(single_loss),
    in_axes=(None, 0, 0),
)

含义是:

  1. 对一个样本的损失求梯度;
  2. 再将这个梯度函数沿 batch 轴映射。

适合逐样本梯度。

7.4 grad(vmap(f))

如果 vmap(f) 输出逐样本结果,不能直接使用普通 grad,因为结果不是标量。通常要先聚合:

def mean_batch_loss(params, x_batch, y_batch):
    losses = jax.vmap(single_loss, in_axes=(None, 0, 0))(
        params,
        x_batch,
        y_batch,
    )
    return jnp.mean(losses)

batch_grad = jax.grad(mean_batch_loss)

顺序不是形式上的偏好,而是由函数的输入输出类型决定的。

7.5 pmapvmap

可以在单设备内部使用 vmap,再用 pmap 分布到多个设备:

pmap(
    vmap(
        单样本计算
    )
)

这表示:

  • pmap 处理设备维;
  • vmap 处理设备内 batch 维。

例如数据形状可能是:

(num_devices, local_batch, feature_dim)

但实际是否需要显式 vmap,取决于单设备函数是否已经使用批量矩阵运算。对于 x @ w 这种操作,JAX 本身已经能处理 batch 维,不一定需要再包一层 vmap


八、设备执行:数组在哪里,计算何时完成

8.1 查看设备和后端

import jax

print(jax.devices())
print(jax.local_devices())
print(jax.default_backend())

可能看到:

[CpuDevice(id=0)]
cpu

也可能是 CUDA GPU 或 TPU。具体输出取决于安装的 JAX 版本、平台插件和运行环境。

JAX 数组通常会被放置到默认后端,也可以显式放置:

import jax.numpy as jnp

x = jnp.ones((1024, 1024))
print(x.devices())

不同版本中 jax.Array 的设备查询接口和 sharding 展示可能有所差异,因此生产代码不应依赖某个调试字符串的精确格式。应关注数组是否被放置在预期设备,以及跨设备传输是否发生。

8.2 主机数组与设备数组

下面的 NumPy 数组位于主机内存:

import numpy as np

x_host = np.ones((1024, 1024), dtype=np.float32)

传入 JAX 运算后,JAX 可能把它复制到设备:

import jax.numpy as jnp

x_device = jnp.asarray(x_host)

如果训练循环每一步都从主机创建或复制 batch,就可能产生主机到设备的传输开销。通常应:

  • 提前准备或预取 batch;
  • 保持数据 dtype 稳定;
  • 避免在 JIT 函数内部调用 NumPy 或 Python 数据处理;
  • 避免频繁将设备数组转回 NumPy。

8.3 异步派发

设备计算常常是异步提交的:

y = model(x)
print("submitted")

打印 "submitted" 只说明 Python 线程已经提交计算,不代表设备已经完成。

如果要测量实际执行时间,应调用:

y.block_until_ready()

完整测量示例:

import time
import jax
import jax.numpy as jnp

@jax.jit
def compute(x):
    return jnp.sin(x) @ jnp.sin(x).T

x = jnp.ones((2048, 128))

# 预热,排除首次编译时间
compute(x).block_until_ready()

start = time.perf_counter()
y = compute(x)
y.block_until_ready()
elapsed = time.perf_counter() - start

print(f"execution time: {elapsed:.4f}s")

如果缺少 block_until_ready(),计时结果可能只反映 Python 提交操作的速度,而不是设备真正执行完成的时间。

8.4 何时必须同步

以下场景通常需要显式同步:

  • 性能基准测试;
  • try 块中确认设备计算错误;
  • 设备计算结果即将用于主机逻辑;
  • 训练步骤结束后记录准确指标;
  • 保存检查点前确保相关计算完成。

例如:

try:
    loss = jitted_step(...)
    loss.block_until_ready()
except Exception as exc:
    print("device computation failed:", exc)
    raise

异步执行会改变故障出现的位置:错误可能在提交调用之后才暴露。因此日志中的“步骤函数已返回”不一定意味着该步骤成功完成。


九、JAX 计算中的错误路径和诊断方法

9.1 形状不一致

jit 通常会针对不同形状生成不同编译版本。若模型期望:

x.shape = (batch, 4)
w.shape = (4, 2)

而实际输入是 (batch, 5),会在矩阵乘法处失败。

建议在进入训练步前明确检查数据契约:

def validate_batch(x, y):
    if x.ndim != 2:
        raise ValueError(f"x must be rank 2, got {x.shape}")
    if y.ndim != 2:
        raise ValueError(f"y must be rank 2, got {y.shape}")
    if x.shape[0] != y.shape[0]:
        raise ValueError(
            f"batch mismatch: x={x.shape}, y={y.shape}"
        )

检查可以放在 JIT 外部,因为它属于输入验证,而不是设备数值计算。

9.2 JIT 重编译

一个常见问题是输入形状不断变化:

for batch in data_loader:
    loss = jit_train_step(params, batch.x, batch.y, 0.01)

如果最后一个 batch 的大小不同,可能触发新的编译。序列长度变化也会产生类似问题。

诊断时可以:

  • 观察日志中的编译警告;
  • 临时开启 JAX 的编译日志;
  • 统计不同输入形状;
  • 使用固定 batch size;
  • 对变长序列进行 padding 或 bucketing。

固定形状不是数学上的必要条件,但常常是稳定编译缓存和推理延迟的重要条件。

9.3 NaN、Inf 和数值不稳定

grad 能正确计算并不保证训练稳定。以下情况可能产生 NaNInf

  • 学习率过大;
  • 指数、对数、除法的输入越界;
  • 混合精度下溢或溢出;
  • 梯度爆炸;
  • 归一化分母接近零。

可以在调试阶段检查:

def check_tree_finite(tree):
    leaves = jax.tree_util.tree_leaves(tree)
    return all(bool(jnp.all(jnp.isfinite(x))) for x in leaves)

但这个函数中把设备数组转换为 Python bool,不应放进 JIT 训练步作为正常计算逻辑。可以在同步后进行检查:

params, loss = jit_train_step(params, x, y, 0.01)
loss.block_until_ready()

if not bool(jnp.isfinite(loss)):
    raise FloatingPointError(f"non-finite loss: {loss}")

生产系统还应保存最近的参数、随机 key、数据批次标识和配置,以便从异常步骤恢复或复现。

9.4 变换内部的调试打印

普通 print 位于 JIT 函数内部时,往往只在追踪或编译阶段打印,而不是每次执行都打印。需要设备执行时的调试输出,可使用:

jax.debug.print("loss = {}", loss)

它适合短期调试,但会引入同步或通信成本,不应无条件保留在高频训练路径中。


十、一个端到端的组合示例

下面将纯函数、gradjitvmap 和设备执行放在一个小型训练流程中。

import time
import jax
import jax.numpy as jnp

def predict(params, x):
    return x @ params["w"] + params["b"]

def single_loss(params, x, y):
    pred = predict(params, x)
    return jnp.mean((pred - y) ** 2)

def batch_loss(params, x_batch, y_batch):
    # 先得到每个样本的损失,再求批量平均
    per_example_losses = jax.vmap(
        single_loss,
        in_axes=(None, 0, 0),
    )(params, x_batch, y_batch)
    return jnp.mean(per_example_losses)

def update(params, x_batch, y_batch, learning_rate):
    loss, grads = jax.value_and_grad(batch_loss)(
        params,
        x_batch,
        y_batch,
    )

    new_params = jax.tree_util.tree_map(
        lambda p, g: p - learning_rate * g,
        params,
        grads,
    )
    return new_params, loss

jit_update = jax.jit(update)

key = jax.random.key(7)
key, kx, kw = jax.random.split(key, 3)

batch_size = 16
feature_dim = 3
output_dim = 1

x = jax.random.normal(kx, (batch_size, feature_dim))

true_w = jnp.array([[2.0], [-1.0], [0.5]])
true_b = jnp.array([0.2])
y = x @ true_w + true_b

params = {
    "w": jax.random.normal(kw, (feature_dim, output_dim)),
    "b": jnp.zeros((output_dim,)),
}

# 首次调用包含编译
params, loss = jit_update(params, x, y, 0.05)
loss.block_until_ready()

start = time.perf_counter()

for _ in range(20):
    params, loss = jit_update(params, x, y, 0.05)

# 等待最后一次设备执行完成后再计时
loss.block_until_ready()
elapsed = time.perf_counter() - start

print("final loss:", float(loss))
print("elapsed:", elapsed)

执行路径是:

  1. single_loss 定义单样本损失;
  2. vmap(single_loss) 将其提升为批量损失;
  3. batch_loss 对样本损失求平均,得到标量;
  4. value_and_grad(batch_loss) 计算损失和参数梯度;
  5. update 使用显式参数树完成 SGD 更新;
  6. jit_update 将整个训练步编译;
  7. block_until_ready 确保计时和结果观察对应真实设备完成状态。

这个结构也适合扩展到深度网络。需要增加的只是:

  • 更复杂的参数 PyTree;
  • 显式的优化器状态;
  • 训练和评估函数;
  • 随机 key;
  • 设备分片或复制逻辑。

十一、模型、数据、评测与生产边界

11.1 训练状态必须区分参数和非参数状态

实际模型除了参数,还可能有:

  • 优化器动量;
  • BatchNorm 统计量;
  • 学习率调度器状态;
  • 随机数 key;
  • 训练步数;
  • 混合精度损失缩放器。

这些对象不能都混成一个“参数数组”。更准确的状态接口类似:

state = {
    params,
    optimizer_state,
    model_state,
    rng_key,
    step,
}

训练函数显式接收并返回完整状态:

(state, batch) -> (new_state, metrics)

评测函数则应明确使用评测状态,避免把验证数据、权限配置或线上数据写入训练状态。

11.2 变换不负责权限隔离

jitvmappmap 只理解数组计算,不理解:

  • 用户权限;
  • 数据租户;
  • 训练集和测试集的业务归属;
  • 个人信息脱敏;
  • 模型是否允许访问某个数据源。

如果把数据复制到多个设备,pmap 会在设备间复制或传输这些数据。敏感数据的权限边界必须在进入 JAX 计算前建立,并在检查点、缓存、日志和设备内存生命周期中持续管理。

11.3 设备副本会增加资源成本

pmap 的复制式数据并行通常意味着每个设备保存一份参数副本。参数、梯度、优化器状态和激活值都会影响显存或内存占用。

在生成式 AI 模型中,优化器状态可能比参数本身还大。多设备并行并不自动降低总内存使用;它主要把计算和局部内存压力分散到多个设备。若目标是模型参数分片,而不是完整复制,应使用适合的 sharding 方案。

11.4 评测要避免无意中的数据并行语义错误

例如在多设备评测中,每个设备计算局部准确率:

ad=correctdcountda_d = \frac{\text{correct}_d}{\text{count}_d}

若各设备样本数相同,可以对准确率做 pmean。若样本数不同,则应聚合计数:

a=dcorrectddcountda = \frac{\sum_d \text{correct}_d} {\sum_d \text{count}_d}

错误地对局部准确率做简单平均,会在最后一个 batch 或过滤样本后产生偏差。这是 pmap 集体通信中常见的“形状正确但统计含义错误”。


十二、几个需要明确区分的结论

12.1 规范保证与实现表现

可以相对确定地依赖:

  • JAX 数组具有不可变语义;
  • grad 对标量目标返回梯度;
  • vmap 沿指定轴执行批量映射;
  • pmap 在映射轴上的设备之间支持命名集体通信;
  • jit 可以将 JAX 数值计算编译到后端。

不应把以下内容当作永久保证:

  • 某个具体操作一定融合成一个内核;
  • 某个数组一定驻留在某一块设备;
  • print 在 JIT 内部一定按调用次数执行;
  • 不同形状调用一定共享编译结果;
  • pmap 在所有未来版本中都是首选的多设备 API。

12.2 jit 不等于一定更快

如果函数很小,编译成本可能超过执行收益。以下情况尤其需要谨慎:

  • 只调用一次的短函数;
  • 输入形状频繁变化;
  • 函数内部包含大量 Python 控制逻辑;
  • 每一步都发生主机设备拷贝;
  • 计算本身很小,但频繁同步。

应使用预热后的、包含 block_until_ready() 的测量结果判断收益,而不是只看 Python 函数返回时间。

12.3 vmap 不等于复制 Python 循环

vmap 是对 JAX 运算的程序变换,不是简单地把 Python 循环展开 NN 次。它需要函数中的操作具有相应的批量规则。对于没有 JAX batching rule 的自定义操作,可能无法直接向量化,需要改写操作或注册规则。

12.4 pmap 不等于自动同步一切状态

只有通过 lax.pmeanlax.psum 等集体操作显式聚合的数据才会同步。每个设备上的 Python 状态、外部文件操作和独立随机 key 不会因为使用 pmap 就自动保持一致。


结语:把 JAX 程序写成可变换的数值函数

JAX 工程的核心不是分别记忆四个 API,而是建立如下思维模型:

纯函数
  ├─ grad:改变函数,得到导数函数
  ├─ vmap:改变函数,得到批量函数
  ├─ jit:改变函数,得到编译执行函数
  └─ pmap:改变函数,得到跨设备映射函数

当参数、随机 key、优化器状态、模型状态和数据都通过输入输出显式传递时,训练函数可以被组合、编译、向量化和复制到设备。反之,如果计算依赖隐藏状态、动态 Python 分支、隐式随机性或未控制的主机设备传输,问题通常不会在最初的函数定义处暴露,而会在追踪、重编译、多设备同步、故障恢复或评测统计阶段出现。

因此,一个可靠的 JAX 训练步应尽量接近:

(state,batch)(new_state,metrics)(\text{state}, \text{batch}) \longrightarrow (\text{new\_state}, \text{metrics})

jit 决定它如何编译,grad 决定它如何获得优化方向,vmap 决定它如何批量化,pmap 决定它如何跨设备执行,而纯函数决定这些变换能否保持清晰、可验证和可恢复。


系列导航与关联阅读

官方资料

本文依据研究论文、标准组织与主流框架官方文档重新梳理;正文、示例与工程清单由 WR BLOG 编写。