You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.30 07:27:52