
本文介绍在不拆分批次的前提下,利用 jax.vmap + jax.jacfwd 对批量图像函数计算逐样本雅可比矩阵,并通过张量重塑与对角线求和高效提取散度(即雅可比矩阵的迹),适用于神经网络输出层的可微几何分析任务。
本文介绍在不拆分批次的前提下,利用 `jax.vmap` + `jax.jacfwd` 对批量图像函数计算逐样本雅可比矩阵,并通过张量重塑与对角线求和高效提取散度(即雅可比矩阵的迹),适用于神经网络输出层的可微几何分析任务。
在深度学习中,常需对神经网络输出(如 Flow 模型、Normalizing Flow 或扩散模型中的变换)计算逐样本雅可比矩阵的散度(divergence)——即每个样本输出关于其自身输入的雅可比矩阵的迹(trace)。但 jax.jacfwd(f) 默认将整个输入张量视为一个扁平向量,当 f 接收形状为 [k, H, W, C] 的批量图像时,其雅可比会跨样本耦合,产生冗余的 [k, H, W, C, k, H, W, C] 形状,而我们真正需要的是 每个样本独立的 [H, W, C, H, W, C] 雅可比子块,最终目标是得到长度为 k 的散度向量。
核心思路是:不改变 f 的批量接口,而是用 jax.vmap 将 jacfwd 的语义从“全批量联合微分”转变为“沿 batch 维度并行执行单样本微分”。这既保持了模型前向推理的批处理效率,又避免了手动循环调用导致的性能损失。
✅ 正确做法:vmap 包裹 jacfwd
首先,确保你的函数 f 在单个样本输入下定义良好(即支持形状为 [H, W, C] 的输入)。即使原始 f 是为批量设计的,也应重构为接受单样本,并由 vmap 自动广播:
import jax
import jax.numpy as jnp
def f_single(x): # x: [H, W, C]
# 示例:逐元素平方(可替换为 model.apply(params, x))
return x * x, x * x # 返回 (output, aux)
# 关键:先对单样本函数求雅可比,再沿 batch 维度向量化
f_batch_jac = jax.vmap(
jax.jacfwd(f_single, has_aux=True),
in_axes=0, # 沿第 0 维(batch)映射
out_axes=(0, 0) # 雅可比与辅助值均保留 batch 维
)
此时,若输入 x 形状为 [k, H, W, C],则 f_batch_jac(x) 输出:
- jac: 形状为 [k, H, W, C, H, W, C] —— 每个样本对应一个完整雅可比矩阵
- val: 形状为 [k, H, W, C] —— 与原需求一致,且 has_aux=True 保证了值的同步返回
? 计算散度(Jacobian 的迹)
散度定义为:对每个样本,将其雅可比矩阵(形状 [H, W, C, H, W, C])视为二维矩阵([D, D],其中 D = H*W*C),然后取对角线元素之和。
推荐做法是避免显式展开高维数组,而使用 einsum 实现清晰、高效的迹计算:
def divergence(jac):
# jac: [k, H, W, C, H, W, C]
k, H, W, C, _, _, _ = jac.shape
D = H * W * C
# 展平空间+通道维度:[k, D, D]
jac_flat = jnp.reshape(jac, (k, D, D))
# 计算迹:对每个 k,求 jac_flat[i, j, j] 的和 → [k]
return jnp.einsum('kii->k', jac_flat)
# 完整流程
x = jnp.ones((3, 2, 2, 1)) # 示例输入
jac, val = f_batch_jac(x)
div = divergence(jac) # shape: (3,)
print("Divergence:", div) # e.g., [4. 4. 4.] for x=1
⚠️ 注意事项:
- vmap 要求 f_single 是纯函数(无副作用、确定性),且所有控制流(如 for 循环)必须能被 JAX 追踪;建议用 jnp.arange + jnp.where 或 lax.fori_loop 替代 Python 循环。
- 若 f_single 内部调用大型神经网络(如 model.apply),请确保 params 已通过闭包或 static_argnums 正确传入,且模型本身支持单样本前向。
- einsum('kii->k') 比 jnp.trace(jac_flat, axis1=1, axis2=2) 更直观且在 JIT 下通常更优。
? 进阶优化:直接计算散度(无需存储完整雅可比)
若内存受限,可使用 jax.jacrev + vmap 结合 jnp.sum 实现前向模式散度(Forward-mode divergence),跳过显式构造雅可比:
def divergence_direct(f_single, x):
def trace_fn(x_i):
y, _ = f_single(x_i)
# 使用 forward-mode:对每个输入基向量方向求导并累加 dy_i/dx_i
D = x_i.size
eye = jnp.eye(D).reshape(D, *x_i.shape)
# vmap over basis vectors
grads = jax.vmap(lambda e: jax.jvp(lambda xi: f_single(xi)[0], (x_i,), (e,))[1])(eye)
# grads: [D, *y.shape] → sum diagonal of Jacobian: sum(grads[i, i, ...])
return jnp.sum(jnp.diag(grads.reshape(D, D)))
return jax.vmap(trace_fn)(x)
但该方法复杂度更高,实践中推荐首选 vmap + jacfwd + einsum 方案——它简洁、可读性强、易于调试,且在 GPU/TPU 上经过充分优化。
综上,vmap(jacfwd(...)) 是解决“批量函数逐样本雅可比”问题的标准范式;配合 einsum 迹计算,即可高效、准确地获得每张图像变换的散度,为密度估计、正则化或几何约束提供关键梯度信号。











