JAX中高效计算指定高阶偏导数的方法咨询
问题:如何仅计算三阶偏导的指定分量并避免冗余计算?
已知用于计算函数导数的代码如下:
import jax import jax.numpy as jnp def f(x): return jnp.prod(x) df1 = jax.grad(f) df2 = jax.jacobian(df1) df3 = jax.jacobian(df2)
通过上述代码可获取所有偏导数,结合vmap能得到指定分量:
x = jnp.array([[ 1., 2., 3., 4., 5.], [ 6., 7., 8., 9., 10.], [11., 12., 13., 14., 15.], [16., 17., 18., 19., 20.], [21., 22., 23., 24., 25.], [26., 27., 28., 29., 30.]]) df3_x0_x2_x4 = jax.vmap(df3)(x)[:, 0, 2, 4] print(df3_x0_x2_x4) # [ 8. 63. 168. 323. 528. 783.]
现需要仅计算df3_x0_x2_x4这一分量,避免不必要的导数计算,且保持f为单向量参数。
解决方案
核心思路是逐步求导并仅提取所需分量,避免计算完整的高阶导数张量。三阶偏导∂³f/∂x₀∂x₂∂x₄可以通过三次嵌套的jax.grad实现,每次求导后只保留目标维度的分量,这样每一步都只计算必要的导数信息,大幅减少计算量。
具体代码实现如下:
import jax import jax.numpy as jnp def f(x): return jnp.prod(x) # 定义仅计算目标三阶偏导分量的函数 def compute_target_derivative(x): # 第一步:计算f对x₄的偏导 df_dx4 = lambda x: jax.grad(f)(x)[4] # 第二步:对上述结果求x₂的偏导 d2f_dx2dx4 = lambda x: jax.grad(df_dx4)(x)[2] # 第三步:对上述结果求x₀的偏导,得到目标分量 return jax.grad(d2f_dx2dx4)(x)[0] # 用vmap处理批量输入 x = jnp.array([[ 1., 2., 3., 4., 5.], [ 6., 7., 8., 9., 10.], [11., 12., 13., 14., 15.], [16., 17., 18., 19., 20.], [21., 22., 23., 24., 25.], [26., 27., 28., 29., 30.]]) df3_x0_x2_x4 = jax.vmap(compute_target_derivative)(x) print(df3_x0_x2_x4) # [ 8. 63. 168. 323. 528. 783.]
原理说明
- 每次调用
jax.grad时,仅对前一步的标量结果(目标分量)求导,而非对整个向量/张量求导,因此不会产生冗余的导数计算。 - 嵌套的
grad调用严格对应三阶偏导的求解顺序,最终得到的就是∂³f/∂x₀∂x₂∂x₄的结果。 - 结合
vmap可以高效处理批量输入,和原方法结果完全一致,但计算效率更高。
内容的提问来源于stack exchange,提问作者mrkwjc
相关产品推荐
相关产品推荐

