JAX函数变换深度解析:grad/jit/vmap/pmap组合实战与踩坑指南
2026/9/9 20:29:00 网站建设 项目流程

做深度学习的人第一次打开JAX,多半是被它那句“像NumPy一样写代码,却自带自动微分”给吸引过来的。但我用得越久越觉得,自动微分只是JAX的入场券,真正让它和PyTorch、TensorFlow拉开差距的,是整套函数变换(function transformation)的设计哲学。你可以把JAX理解成一个“函数的加工厂”:普通的函数丢进去,出来的是带梯度、被编译、能向量化、能并行的新函数。这种思路在传统框架里几乎看不到。

这篇文章想聊的,就是JAX函数变换的深层机制和我在实际项目中用到的高阶玩法。我会从JAX最核心的设计讲起,把grad、jit、vmap、pmap逐个拆开,再落到几个组合使用的实战案例上,最后聊聊我自己在工程里踩过的坑。内容主要面向已经会用PyTorch或NumPy、想深入理解JAX的读者,当然,如果你是刚入门的新手,前两个章节也足够帮你建立起正确的认知框架。

1. JAX的真正核心:不只是自动微分

1.1 从自动微分到函数变换

JAX最早的定位是“可微分的NumPy”,很多教程也把重点放在grad上。但如果只把JAX当成一个梯度计算工具,你会错过它最值钱的部分。JAX的核心抽象其实是函数变换——输入一个函数,返回一个函数,中间改变的是这个函数的行为,而不是用户的数据或模型结构。

自动微分只是函数变换里最常见的一种。jax.grad(f)做的事情是:给定函数f,构造一个新函数g,使得g(x)返回f在x处的梯度。这个过程不是把公式写死,而是通过追踪执行轨迹、反向传播两个阶段动态完成的。类似的,jax.jit(f)把f编译成XLA(Accelerated Linear Algebra)执行,jax.vmap(f)把f向量化,jax.pmap(f)把f并行化到多设备。

这种设计带来一个直接好处:你可以自由组合这些变换。比如jax.jit(jax.grad(f))意思是“先算梯度,再编译整个梯度计算过程”;jax.vmap(jax.grad(f))意思是“对每个样本分别求梯度”。在传统框架里,“对batch中每个样本单独算梯度”和“一次算完整个batch的梯度”是两套完全不同的实现,但在JAX里只是函数变换的嵌套组合。

还有个容易忽略的点:函数变换要求函数本身是纯函数(pure function),也就是输出只依赖输入,不读取或修改全局状态。这个约束换来了极大的灵活性——因为纯函数没有隐藏依赖,XLA可以放心地重排执行顺序、合并算子、并行调度。PyTorch里想要做类似的事就麻烦得多,因为它的autograd系统绑定在tensor的mutability上,你必须小心翼翼地管理requires_gradno_grad上下文。

1.2 纯函数约束与函数式编程根基

说“纯函数约束”可能有点抽象,我用一个例子来说明。假设你写了这样一个函数:

import jax.numpy as jnp count = 0 def f(x): global count count += 1 return x * 2

这个函数修改了全局变量count,它不是纯函数。如果你对f调用jax.grad,JAX会直接报错或者给出不可预期的结果。因为梯度计算需要重新执行f的前向过程,而全局状态的变化会让执行轨迹和输入之间失去严格对应关系。

JAX体系下正确的做法是:

def f(x, count): return x * 2, count + 1

把状态显式地作为输入传进去,再把更新后的状态作为返回值传出来。这其实就是函数式编程里“状态即数据”的思想。习惯了PyTorch的in-place操作后,刚转JAX的人总觉得束手束脚,但适应之后会发现,这种约束让代码的可测试性、可组合性都好了很多。

我在实际工程里对这个约束的体会尤其深。一次我在写强化学习算法时,需要维护一个环境状态,最初把状态放在类属性里,结果调用jax.jit后怎么都不对。排查了很久才发现是状态被隐式修改导致追踪失败。改成显式传入传出状态后,代码不仅跑通了,而且因为jit编译了整个训练step,训练速度直接提升了近三倍。

2. 四大函数变换逐个拆解

2.1 grad:梯度从哪来,怎么用

jax.grad是大家最熟悉的入口。它的核心机制是Reverse-mode自动微分(也就是反向传播),但实现方式和PyTorch有很大区别。PyTorch在每次前向计算时构建一个动态计算图,反向传播时按图回传梯度;JAX则通过tracer(追踪器)记录操作序列,然后用XLA编译成高效的梯度计算程序。

基本用法很简单:

import jax import jax.numpy as jnp def loss_fn(w, x, y): pred = jnp.dot(x, w) return jnp.mean((pred - y) ** 2) w = jnp.ones(3) x = jnp.array([[1.0, 2.0, 3.0], [4.0, 5.0, 6.0]]) y = jnp.array([1.0, 2.0]) grad_fn = jax.grad(loss_fn) grads = grad_fn(w, x, y)

jax.grad默认对第一个参数求导,如果想对多个参数求导数,可以指定argnums

grad_fn = jax.grad(loss_fn, argnums=(0, 1)) g_w, g_x = grad_fn(w, x, y)

这里g_w是loss对w的梯度,g_x是loss对x的梯度。argnums这个参数非常实用,特别是在处理多组可学习参数时,不用把参数打包成一个巨型tensor再切片。

需要注意的一点是,jax.grad要求函数输出是标量。如果你的函数返回的是向量或矩阵,需要先通过求和、求均值等方式降维,否则会报错。这是很多人第一次用JAX时遇到的报错之一。

高阶导数也是grad的强项。因为grad本身返回一个函数,你可以对它再次作用grad

def f(x): return jnp.sum(x ** 3) f_prime = jax.grad(f) # 一阶导 f_double_prime = jax.grad(f_prime) # 二阶导 print(f_double_prime(2.0)) # 12.0

这个能力在物理模拟、几何处理里特别有用,比如计算Hessian矩阵或者曲率信息。PyTorch虽然也能做高阶导数,但实现起来要小心处理create_graph=True,相比之下JAX的写法简洁得多。

2.2 jit:即时编译的加速原理与适用边界

jax.jit通过XLA把Python函数编译成高效的融合算子。编译后的函数不再是逐行解释执行Python字节码,而是直接在底层执行,并且多个操作会被融合到一个kernel里,减少访存开销。

使用方法:

@jax.jit def predict(w, x): return jnp.dot(x, w) # 或者 predict = jax.jit(predict)

jit的加速效果在复杂计算图上非常明显。一个包含几十个矩阵乘法、激活函数、归一化操作的神经网络层,如果全部用jit编译,能比纯Python执行快一倍以上。我做过一个简单的基准测试,同样一个MLP的前向传播,在CPU上jit之后大概能快1.5-2倍,在GPU上差距更明显,因为kernel融合减少了大量的CPU-GPU数据传输。

不过jit不是万能的。它有几个“只可意会不可言传”的限制。最重要的是:被编译函数内部不能有Python的side effect,不能依赖外部全局变量,不能用动态Python控制流(比如if x > 0这种依赖运行期数值的条件判断)。如果函数内部有这类操作,XLA在追踪时无法确定执行路径,会报错。

这个限制其实有变通方法。对于Python控制流,JAX提供了结构化原语jax.lax.condjax.lax.while_loop,它们是XLA能够处理的“函数式控制流”。另外,jit函数接受的输入shape和dtype必须是固定的,一旦输入shape变了,JAX会重新编译,造成巨大的编译开销。所以实际工程里要特别注意避免动态shape。

专门说一个常见误区:很多人以为jit只要在函数上加上装饰器就会自动生效,但实际上每次调用时如果传入的Python对象类型不一致(比如第一次传list,第二次传numpy array,第三次传jnp array),JAX会因为无法缓存编译结果而反复重编译。建议所有传给jit函数的数据统一使用jnp.ndarray类型,并且固定shape。

2.3 vmap:自动向量化,告别手工batch

jax.vmap是我认为JAX最被低估的变换。它解决的是“批量计算”问题。在NumPy或PyTorch时代,要给一个函数增加batch维度,通常要改写函数内部的所有运算,加上一个维度并调整broadcast规则。非常繁琐,而且容易出错。

vmap的做法不同:它自动把函数映射到batch维度上。你只需要写针对单个样本的逻辑,然后用vmap一包,就能批量处理任意batch size的数据,而且底层是高效的向量化计算,不是简单的Python循环。

def single_kernel(x): return jnp.sin(x) + x ** 2 batch_kernel = jax.vmap(single_kernel) x_batch = jnp.arange(10.0) print(batch_kernel(x_batch))

更复杂的场景是多个参数带不同的batch维度。比如在对比学习里,query和key分别来自两个batch,你需要对batch内的每个query和所有key计算相似度矩阵。这时可以组合两个vmap

def sim(q, k): return jnp.dot(q, k) # 第一个vmap把q映射到batch维度 # 第二个vmap把k映射到batch维度,整体得到一个batch x batch的相似度矩阵 similarity_matrix = jax.vmap(lambda q: jax.vmap(lambda k: sim(q, k))(keys))(queries)

写这行代码前,我建议先在纸上画一下维度的变化。jax.vmapin_axes参数可以指定哪些输入参与batch映射,默认所有输入都在第0维。上面这个例子里,keys没有跟着query的batch走,所以需要用嵌套vmap把k的维度“内层化”。

vmap还有一个容易被忽视的优势:它通常比手动写batch循环更快。因为vmap生成的计算图会经XLA编译,循环展开和算子融合都是由编译器完成的,即便你写的是Python层的“逻辑循环”,编译后也是高效的向量化计算。这让我在实现各类元学习算法时省了很多事,不需要为每个batch变体单独维护一套代码。

2.4 pmap:多设备并行与分布式扩展

jax.pmap用于把计算分布到多个设备(GPU/TPU)上。它和vmap很像,也是沿某个轴“映射”,但映射的单位是设备而不是单条数据。调用pmap后,JAX会把你定义的函数复制到每个设备上,每个设备处理输入的一个切片,通过设备间通信协议同步结果。

典型用法:

jax.pmap(lambda x: x * 2)(jnp.arange(8.0).reshape(2, 4))

上面这个例子把8个元素分成两组,两个设备各处理一组。pmap的返回值在所有设备上的结果会通过all-gather方式汇聚回来。

实际工程里,pmap最常见的场景是数据并行训练。模型参数在每个设备上复制一份,每个设备处理一个mini-batch,然后通过jax.lax.pmean等集合通信操作同步梯度。

写并行训练代码时最容易踩的坑是:pmap里的函数如果有随机数生成,每个设备必须使用不同的随机key。JAX的随机数设计是显式传入key的,这反而帮了大忙——只要每个设备的key不同,随机性天然就是per-device的,不会出现所有设备生成同样随机数的情况。

3. 高阶应用实战:组合变换

3.1 组合grad+jit+vmap:一个完整的深度学习训练闭环

函数变换最大的魅力在于可以任意组合。下面我写一个小型逻辑回归的完整训练流程,展示grad、jit、vmap是如何协同工作的。

import jax import jax.numpy as jnp def model(params, x): w, b = params return jnp.sigmoid(jnp.dot(x, w) + b) def loss_fn(params, x, y): pred = model(params, x) return -jnp.mean(y * jnp.log(pred + 1e-7) + (1 - y) * jnp.log(1 - pred + 1e-7)) @jax.jit def train_step(params, x, y): grads = jax.grad(loss_fn)(params, x, y) # 手动更新,展示梯度更新的透明度 new_params = [(w - 0.1 * dw, b - 0.1 * db) for (w, b), (dw, db) in zip(params, grads)] return new_params # 初始化 key = jax.random.PRNGKey(0) key, subkey = jax.random.split(key) w = jax.random.normal(subkey, (2,)) b = jnp.array(0.0) params = (w, b) # 生成数据 X = jax.random.normal(jax.random.PRNGKey(1), (100, 2)) y = (X[:, 0] + X[:, 1] > 0).astype(jnp.float32) for step in range(10): params = train_step(params, X, y) if step % 2 == 0: print(f"step {step}, loss {loss_fn(params, X, y):.4f}")

这里train_step被整体jit编译了,里面依然可以调用jax.grad,因为JAX的变换是嵌套透明的。grad产生的梯度计算图,和loss函数本身的算子,都会被XLA融合到同一个编译单元里。这是JAX能实现“编译整个训练过程”的根源。

如果要处理batch,只需要把数据按batch维度切分,再在train_step外面套一层jax.vmap或者手动写循环。因为本身train_step只处理一个batch的数据,如果数据太大,可以用vmap把“对每个micro-batch更新一步”的逻辑向量化:

def micro_step(params, x_batch, y_batch): return train_step(params, x_batch, y_batch) # 注意:grad的更新是在每个micro-batch上独立进行的 # 现实中往往用grad累积,这里只是展示组合语法 params_per_batch = jax.vmap(lambda x, y: micro_step(params, x, y))(X_batches, y_batches)

在实际项目里,我不建议直接vmap训练step,因为通常我们要的是所有micro-batch梯度累积后再更新一次,而不是每个batch独立更新。这个例子主要说明JAX的变换组合能力有多强——你几乎可以用任何自然的方式拼装逻辑,剩下的性能问题交给编译器。

3.2 自定义变换:从grad到jacfwd、jacrev

除了grad,JAX还内置了求Jacobian矩阵的jacfwd(前向模式)和jacrev(反向模式)。这两者都是基于grad机制构建的高阶变换。理解它们的关键在于模式的差异:

  • 反向模式(reverse-mode,即jacrev):适合输入维度远大于输出维度的函数,计算复杂度与输出维度成正比。
  • 前向模式(forward-mode,即jacfwd):适合输出维度远大于输入维度的函数,计算复杂度与输入维度成正比。

举个例子,假设函数f把R^2映射到R^1000,那么jacfwd明显更划算,因为它只需要执行2次前向传播;而jacrev需要执行1000次反向传播。反之,如果函数把R^1000映射到R^2,则jacrev更优。

def f(x): return jnp.array([x[0] ** 2, x[0] * x[1], x[1] ** 3]) J = jax.jacfwd(f)(jnp.array([2.0, 3.0])) print(J)

在实际使用中,我经常用jacfwd计算雅可比矩阵用于机器人运动学分析,或者用jacrev计算神经网络输出对输入的敏感度矩阵。它们返回的不是一组gradient,而是完整的Jacobian矩阵,直接铺开就是整个导数矩阵。

JAX还支持自定义变换,jax.custom_jvpjax.custom_vjp允许你手动指定函数的前向和反向导数规则。这个功能在做科学计算时特别有用——例如有些数值稳定的softmax实现、ODE求解器内部的导数规则,用自动微分直接推导会不稳定,这时你可以绕过黑盒,手动写导函数。我写过一次自定义VJP的经验是:虽然需要额外处理很多细节,但对于数值稳定性要求高的优化问题,回报是切实的,优化器收敛速度明显提升。

3.3 科学计算场景:ODE求解、优化问题

JAX天然适合科学计算。原因是它可以对任意数值程序做自动微分和编译,这在物理模拟、最优控制、计算化学等领域价值巨大。以一个简单的常微分方程求解为例:假设你要解dy/dt = -k*y,除了用scipy.integrate,你还可以用JAX的jax.experimental.ode.odeint

import jax import jax.numpy as jnp from jax.experimental.ode import odeint def dynamics(y, t, k): return -k * y y0 = jnp.array([1.0, 2.0]) t = jnp.linspace(0.0, 1.0, 100) ys = odeint(dynamics, y0, t, k=0.5)

更关键的是:你可以对ODE求解过程直接求导。这意味着你可以把“ODE求解器”当成一个可微层放进神经网络里,做参数辨识或者神经ODE(Neural ODE)。这在传统数值计算里很难实现,因为求解器内部有大量的迭代、自适应步长、停止条件,都是不可微的逻辑。JAX因为记录的是整个计算轨迹,可以通过隐函数求导的方式绕过这些不可微点,给出正确的梯度。

优化问题也是典型的应用场景。JAX提供了jax.scipy.optimize.minimize,虽然功能没有SciPy全面,但好处是整个目标函数可以被jit编译,配合自动微分在GPU上跑大规模并行优化。我曾经用JAX做过一个最优运输问题,目标函数里包含几千个点的距离矩阵计算,用jit+grad组合优化后,速度比SciPy版本快了一个数量级。

4. 常见问题与排查技巧实录

4.1 明明调用了jit,为什么没有加速

这是最常被问的问题。执行jax.jit(f)(x)时如果觉得和没编译差不多,先检查三件事:

  1. 输入类型是否统一:如果第一次传numpy.ndarray,第二次传jnp.ndarray,JAX会为每一种类型单独编译。统一后可以显著提高缓存命中率。
  2. 函数内部是否包含大量Python层逻辑:比如外层有Python的for i in range(len(...))循环,每次迭代都在调用jit函数。这种情况下整个jit函数被反复重启,编译收益会被调用开销抵消。应当把循环放进jit函数内部,或者用jax.lax.scan处理循环。
  3. 是否有动态数组操作:比如jnp.nonzerojnp.where在追踪时产生动态shape,会导致XLA无法静态编译,只能fallback到解释模式。

判断一段代码是否真的被编译了,最直接的方法是加一个print在函数外面:把print放在被jit的函数外面,每次执行都会打印;放在函数内部,第一次调用会打印,之后不打印,因为编译后的代码不会执行Python打印语句。

4.2 动态shape、随机性、不可变数组的坑

动态shape的坑特别隐蔽。看这个例子:

@jax.jit def f(x): mask = x > 0 return x[mask] * 2 # 这里返回的shape取决于x的值,是动态的

这种代码在第一个输入上可能正常运行,但换一个输入,可能立刻报“ConcretizationTypeError”。这是因为JAX在追踪时发现mask是一个动态值,无法静态确定x[mask]的长度。解决办法是避免在jit函数内部做基于值的索引/切片,改用定长mask,或者用jnp.where把非法位置填充成0/某个特殊值,保持shape固定。

关于随机性,JAX和PyTorch完全不同。JAX推荐的做法是显式地持有并拆分PRNGKey:

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

如果多个地方需要随机数,就多次split。这个设计的好处是随机性完全可控、可复现;坏处是很多人初学时把同一个key传给了两次jax.random.normal,结果生成了完全相同的两个张量,还以为是bug。实际上这正是“显式随机性”的设计意图。

还有一个反直觉的坑:JAX数组是不可变的。x[0] = 1这种操作直接报错。你会觉得这很烦,但反过来,正因为不可变,XLA才能做激进的优化。在需要更新数组的场景里,用x.at[0].set(1)代替,它会返回一个新的数组,原来的数组保持不变。

4.3 实战经验速查表

下面这张表是我在实际项目中总结出来的,每一次踩坑之后追加一条,现在分享出来:

问题现象根本原因解决方案
代码加了@jax.jit反而更慢Python层循环导致jit频繁re-trigger把循环放进jit内部,或用jax.lax.scan替代
报了ConcretizationTypeError函数内部依赖运行时值做shape判断改用jnp.where或定长mask
两次随机数生成结果一样同一个PRNGKey被重复使用每次生成前calljax.random.split
x[0] = 1直接报错JAX数组不可变x.at[0].set(1)
grad对向量输出报错grad要求输出是标量先求和或mean
pmap训练loss不收敛不同设备上梯度没有同步jax.lax.pmean做all-reduce
手动写了batch循环很慢Python循环性能差jax.vmap替换
使用列表当参数传入jit函数JAX无法缓存不同结构参数统一转成jnp.ndarray

4.4 关于复变函数与积分变换的一点联想

这篇文章写到这,可能有人会觉得“函数变换”这个概念有点抽象。其实如果你学过复变函数与积分变换(很多理工科同学应该对那本教材有印象),就会觉得JAX的设计并不陌生。复变函数里的傅里叶变换、拉普拉斯变换,本质上就是对“函数”这个对象做某种操作,把时域函数变成频域函数,核心思想是“换个视角看问题”。JAX的函数变换也是类似的——不过它变换的是计算行为本身。

我在入门JAX的时候,恰好手边有一本复变函数与积分变换第六版的PDF做参考。虽然那本书讲的是数学变换,不是编程,但里面的思维方式给了我很大启发:很多看似复杂的计算问题,只要选择一个合适的变换角度,就能化繁为简。JAX的vmap就是很好的例子——你不需要反复手写batch逻辑,只需要“变换”一下函数的执行方式,就能让它在任意batch尺寸下高效运行。

这一点我建议所有学JAX的人都多想一想:当你面对一个复杂的计算任务时,先别急着写实现,先问自己——能不能用某种“变换”把这个任务从根源上简化?这种思维方式一旦建立,JAX就不再是一个库,而是一种解题工具。

5. 从组合变换到我的工程体会

5.1 为什么说组合变换是JAX的护城河

前面提到的grad、jit、vmap、pmap,单独拎出来任何一个,别的主流框架也都能找到类似的功能。PyTorch有autograd、torch.compile、broadcasting和DistributedDataParallel,功能上未必差多少。但JAX真正不可替代的,是这些变换可以自由组合,而且组合后的效果是乘性的。

举一个很直观的例子。在实现一个基于模型的强化学习算法时,我需要做这么一件事:对一批轨迹预测结果求动作序列的梯度。每次预测需要用ODE求解器;轨迹有batch维度;动作序列有长有短;而且我希望整个计算能编译执行。这个需求写出来是jax.jit(jax.vmap(jax.grad(ode_predict)))的组合,在一个装饰器里就能表达清楚。

换成PyTorch做同样的事,你需要:

  • torchdiffeq这类库处理可微ODE,它是单独的一套体系;
  • 自己实现batch维度的处理;
  • torch.compile优化整个图,但经常因为控制流复杂而失败。

不是说PyTorch做不到,而是做得极其绕。JAX的函数式设计让“完全可微、向量化、高性能”成为默认行为,而不是需要各种补丁才能实现的特例。这种生态上的优势,是JAX在科学计算、元学习、强化学习等前沿研究领域快速站稳脚跟的根本原因。

5.2 项目实战中的几点心得体会

最后说几个我的个人体会,算是对整篇文章的一个收尾。

第一,JAX不是让你丢掉PyTorch。在工业界做大规模落地的项目,PyTorch的生态成熟度、部署链路、社区积累目前仍然有优势。JAX更强的是在研究探索阶段提供一个灵活、快速、可微的实验环境。

第二,不要一上来就猛上pmap和自定义VJP。JAX的学习曲线是“先会用grad/jit/vmap,再过一遍常见坑,再谈高阶组合”。我见过很多新手一上来就写复杂的pmap代码,结果出了问题完全没办法调试。建议先把单设备、单batch的流程跑通,再一步步做向量化和并行化。

第三,JAX的函数式风格会反向塑造你写代码的思路。用了JAX半年后,我发现即使回去写PyTorch,也会倾向于把状态显式化、把副作用最小化,代码可读性和可维护性反而提升了。这一点算是意外收获。

第四,如果你想深入某个方向,一定要学会用jax.debug.printjax.debug.callback来调试。因为它们可以获取运行时的值,而且不会破坏jit编译。早期我调试jit代码时总是靠“print在外层”的办法,效率很低。后来看了官方文档里调试章节,才算真正上手。

JAX不算一个容易上手的框架,但你一旦摸清了它的设计哲学,会发现自己以前写的很多代码,其实都是在和框架本身对抗。而JAX想做的,是把“把计算过程的变换”这件事交还给你——你可以从容地定义自己的计算逻辑,然后像搭积木一样,把梯度、编译、向量化、并行化一层层叠加上去。这种自由度,在传统框架里很难体验得到。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询