在深度学习模型的分析与调试中,有时需要拿到模型输出向量对全部参数的完整偏导数矩阵,也就是雅可比矩阵。常见的训练流程只关心标量损失函数对参数的梯度,反向传播一次即可得到结果;但当输出本身是高维向量、且我们需要逐样本的梯度信息时,问题就变得复杂得多。例如可解释性研究中要分析输入或参数对输出的敏感度,神经切线核需要计算模型输出对参数的雅可比,逐样本梯度裁剪也需要获取每个样本独立的影响。直接对每个输出分量执行一次反向传播,在输出维度和批量大小稍大时就会遭遇严重的性能瓶颈。

因此高效计算批量雅可比矩阵的关键在于理解自动微分的两种模式,并根据输出维度与参数维度的相对大小选择合适的策略。下文先明确数学定义,再分析计算成本,最后给出可运行的代码示例和优化建议。
一、雅可比矩阵与逐样本梯度的关系
设模型函数为 f: R^n → R^m,输入 x 为 n 维向量,输出 y 为 m 维向量,模型可训练参数 θ 包含 p 个标量。输出对参数的雅可比矩阵 J 定义为所有一阶偏导数组成的 m×p 矩阵,其中第 i 行第 j 列元素为 ∂y_i / ∂θ_j。在批量输入场景中,每个样本都有一个独立的雅可比矩阵,形状为 (batch, m, p)。这个矩阵完整刻画了输出对参数的局部线性敏感度。
训练中常用的标量损失 L 对参数的梯度可以看作雅可比矩阵与损失对输出梯度向量的乘积。具体来说,若损失 L 是输出 y 的函数,那么 ∂L/∂θ = (∂L/∂y)^T · J。当损失为标量时,一次反向传播就能算出 ∂L/∂θ。但当我们需要完整的 J 而不同损失函数时,就必须单独计算每个输出分量对参数的梯度,或者使用前向模式自动微分。这也是为什么直接调用反向传播接口获取全部雅可比会非常低效。
逐样本梯度的需求进一步增加了复杂度。普通训练中的梯度是对一个批次取平均后的结果,而逐样本梯度要求保留每个样本独立的梯度信息。比如在差分隐私训练中,需要对每个样本的梯度进行裁剪和噪声注入;在影响函数分析中,需要知道单个样本对模型参数更新的贡献。这些场景下,批量雅可比矩阵的计算负担会更重,需要专门设计高效实现。
二、自动微分两种模式的计算成本权衡
反向模式自动微分是深度学习框架默认采用的方式。它从输出端开始,沿着计算图反向传播伴随值,一次反向传播可以高效地计算一个标量输出对所有输入或参数的梯度。对于标量损失函数来说,这几乎是完美选择,因为无论参数量 p 多大,一次反向传播的时间复杂度与前向传播相当。然而当输出维度 m 大于 1 时,要得到完整的 m×p 雅可比矩阵,就需要对每个输出分量分别执行一次反向传播,总成本变为 m 次前向加反向传播。当 m 较大时,这种重复计算会迅速变得不可接受。
前向模式自动微分则从输入端开始,为每个输入变量携带一个切向方向,一次前向传播可以高效地计算雅可比矩阵与某个方向向量的乘积 J·v。对于输出维度 m 很大、输入或参数维度 p 很小的情况,前向模式只需 p 次传播即可得到完整雅可比矩阵。但深度学习模型的参数量通常非常大,单独使用前向模式并不现实。不过当输出维度 m 远小于参数维度 p 时,反向模式仍然是更好的选择;反之,如果 m 很大而参与微分的参数量较小,前向模式就可能更高效。
对于批量雅可比矩阵,每个样本的计算相互独立,可以利用向量化技术同时计算多个样本的雅可比。主流框架提供了自动向量化接口,例如 JAX 的 jacrev、jacfwd 和 vmap,PyTorch 的 torch.func 模块也提供了类似能力。这些接口会在底层自动选择反向或前向模式,并把循环逻辑转换为高效的批量操作,大幅减少了用户手动编写循环的成本。
三、批量雅可比矩阵的高效实现方法
最直接的实现方式是对每个样本的每个输出分量调用一次反向传播,并把得到的参数梯度保存下来。这种方法虽然逻辑清晰,但存在两个严重问题:一是需要 m×batch 次反向传播,计算量大;二是每次反向传播都要重新构建或保留计算图,显存占用和 Python 循环开销都非常高。下面代码展示了这种朴素实现,仅适用于极小规模的调试场景。
import torch
model = torch.nn.Linear(10, 5)
x = torch.randn(3, 10)
outputs = model(x)
# 朴素方法:对每个样本的每个输出分量分别反向传播
batch_jacobian = []
for i in range(outputs.size(0)):
sample_jac = []
for j in range(outputs.size(1)):
model.zero_grad()
outputs[i, j].backward(retain_graph=True)
grads = [p.grad.detach().clone() for p in model.parameters()]
sample_jac.append(grads)
batch_jacobian.append(sample_jac)
更好的做法是利用框架提供的自动向量化功能。JAX 的 jax.jacrev 可以一次性计算函数对指定参数的完整雅可比矩阵,并且自动处理输出维度。配合 vmap 可以对批量输入进行映射,从而直接得到逐样本雅可比。下面示例中,f 是一个简单的线性函数,jacrev 对参数计算雅可比,vmap 将映射扩展到 batch 维度,最终返回的 jac 字典中每个参数张量形状为 (batch, output_dim, *param_shape)。
import jax
import jax.numpy as jnp
key = jax.random.PRNGKey(0)
params = {
'W': jax.random.normal(key, (10, 5)),
'b': jax.random.normal(key, (5,))
}
def f(params, x):
return jnp.dot(x, params['W']) + params['b']
x = jax.random.normal(key, (3, 10))
jac = jax.jacrev(f, argnums=0)(params, x)
print(jac['W'].shape) # (3, 5, 10, 5)
PyTorch 从 2.0 开始也提供了 torch.func 模块,用法与 JAX 类似。通过 functional_call 将 nn.Module 转换为纯函数,再用 vmap 和 jacrev 组合即可获得逐样本雅可比。这种方式比手动循环快得多,而且能自动利用 GPU 并行。不过需要注意,不同版本对参数 dict 的支持略有差异,建议在最新稳定版中运行。
import torch
from torch.func import functional_call, vmap, jacrev
model = torch.nn.Linear(10, 5)
params = dict(model.named_parameters())
x = torch.randn(3, 10)
def f(params, x):
return functional_call(model, params, x)
per_sample_jac = vmap(jacrev(f, argnums=0), in_dims=(None, 0))(params, x)
print(per_sample_jac['weight'].shape)
当参数维度极高或输出维度过大导致完整雅可比矩阵无法存储在显存中时,可以考虑低秩近似或随机投影方法。例如使用 Hutchinson 估计器,用随机向量 v 计算 J^T J 的对角线近似,或者通过随机矩阵把高维雅可比投影到低维子空间。这类方法不需要逐分量反向传播,只需要少量加权反向传播即可得到有偏但计算代价很小的估计值,适合神经切线核的谱分析、梯度噪声估计等对精度要求稍宽松的任务。
四、工程实践中的显存与精度优化
计算批量雅可比矩阵时,最大的瓶颈往往不是时间而是显存。完整雅可比张量的形状为 (batch, m, p),对于大规模模型和输出维度可能达到数十亿个元素。为了避免显存溢出,可以采用分块计算策略:每次只对一部分样本或一部分输出分量计算雅可比,然后立即释放中间结果。将计算图分段构建,配合梯度检查点技术,可以在不降低数值精度的前提下将显存峰值降低数倍。
如果输出维度 m 较小但参数维度 p 很大,反向模式每次只返回一个输出分量的梯度,总计算次数为 m,这通常是可以接受的。真正棘手的是 m 和 batch 都很大,同时参数维度也很大。此时可以优先使用前向模式计算 J 与某个低维子空间基的乘积,再用这些乘积重建近似雅可比。或者直接使用框架提供的 jacfwd,在某些模型结构上前向模式反而比多次反向更快。
精度方面,前向模式和反向模式在数学上是精确的,但浮点运算的舍入误差可能导致两种模式得到的结果存在微小差异。一般这不是问题,但在做数值稳定性要求极高的分析时,建议用双精度浮点并对比两种模式的结果。随机投影等近似方法会引入额外误差,需要根据应用场景确定可接受的误差范围。总体原则是:优先使用框架内置的 jacrev 或 vmap 组合,避免手写循环;当显存或时间受限时,再考虑近似策略。
总结来说,高效计算模型输出对所有参数的批量雅可比矩阵,需要根据输出维度、批量大小、参数量和显存预算来选择反向模式、前向模式或近似方法。利用现代框架的自动向量化接口,可以在少量代码中实现高性能计算,避免手动循环带来的性能损失。理解自动微分的底层机制,有助于在工程中做出合理的折中。