如何便捷地对JAX中已被vmap封装的函数求逐行雅可比矩阵?
如何便捷地对JAX中已被vmap封装的函数求逐行雅可比矩阵?
我完全懂你的困扰——手里有个已经用vmap封装好的函数,输入是(M, N)形状的数组,输出也是(M, N),现在要计算逐行的雅可比矩阵得到(M, N, N)的结果,现有的两种方法要么写法hacky不直观,要么内存开销大到没法用,而且还不想动原来的already_vmapped函数对吧?
先简单分析下你现有的两种方法的问题:
- 第一种给输入加dummy维度再嵌套
vmap和jacobian,虽然能得到正确结果,但写法绕来绕去,可读性很差,维护起来也麻烦。 - 第二种直接求整个函数的jacobian再取对角线转置,内存开销爆炸的原因是:
jax.jacobian(already_vmapped, argnums=0)(x, A)会生成一个(M, N, M, N)形状的巨大矩阵,这对稍大一点的M来说内存直接顶不住,完全不实用。
推荐的简洁高效方案
其实核心思路很简单:既然你的already_vmapped函数本身就是对每个样本(行)独立运算的,那我们只需要先定义对单个样本行的求雅可比逻辑,再用vmap批量应用到所有样本上即可——既不用加dummy维度,也不会产生巨大的中间矩阵,写法还非常直观。
代码示例与验证
直接看改进后的写法:
import jax from jax import numpy as jnp rng = jax.random.PRNGKey(42) A = jax.random.normal(rng, shape=(128, 16, 16)) # 你原有的已vmap封装的函数,完全不需要修改 def already_vmapped(x, A): vmap_mul = jax.vmap(jnp.matmul) return vmap_mul(A, x) # 改进后的逐行雅可比计算函数 def clean_row_wise_jac(x, A): # 定义单个样本行的处理逻辑:输入(N,),输出(N,) single_row_fn = lambda x_row: already_vmapped(x_row[None, :], A)[0] # vmap 单样本的雅可比计算到所有行上 return jax.vmap(jax.jacobian(single_row_fn))(x) # 测试验证 x = jnp.ones((128, 16)) print(jnp.allclose(clean_row_wise_jac(x, A), A)) # 输出 True,结果正确
如果你的A参数是每个样本独立的(就像示例里的A是(128,16,16)),也可以很方便地调整vmap的输入轴,让每个样本行对应自己的A矩阵:
def clean_row_wise_jac_per_A(x, A): # 支持每个样本对应独立的A矩阵 single_row_fn = lambda x_row, A_row: already_vmapped(x_row[None, :], A_row[None, :])[0] return jax.vmap( jax.jacobian(single_row_fn, argnums=0), in_axes=(0, 0) # x的每个行对应A的每个矩阵 )(x, A) print(jnp.allclose(clean_row_wise_jac_per_A(x, A), A)) # 同样输出 True
额外小技巧
如果你的输入输出维度差异很大,可以选择用jax.jacfwd(前向模式)或者jax.jacrev(反向模式)来优化效率:
- 当输出维度 ≤ 输入维度时,
jax.jacrev(反向模式)更高效 - 当输入维度 < 输出维度时,
jax.jacfwd(前向模式)更高效
比如换成前向模式的写法:
def clean_row_wise_jac_fwd(x, A): single_row_fn = lambda x_row: already_vmapped(x_row[None, :], A)[0] return jax.vmap(jax.jacfwd(single_row_fn))(x)
结果和之前完全一致,但在特定维度场景下能节省计算时间。
备注:内容来源于stack exchange,提问作者John D
相关产品推荐
相关产品推荐

