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

JAX中custom_jvp返回None输出引发TypeError,求解决方法

JAX自定义JVP时多输出结构不匹配问题的解决办法

你遇到的错误是因为JAX要求自定义JVP规则返回的原输出(primals_out)和切线输出(tangents_out)必须拥有完全一致的PyTree结构,不能用None来表示某个输出不需要微分。

如果想让函数的某个输出不参与微分计算,正确的做法有两种:

方法一:返回匹配类型的零切线值

修改JVP规则,把None换成和原输出同类型的零值,保证结构一致:

from jax import custom_jvp, jacobian

@custom_jvp
def func(x, y):
    return x+y, x*y

@func.defjvp
def func_jvp(primals, tangents):
    x, y = primals
    t0, t1 = tangents

    primals_out = func(x, y)
    # 第二个输出x*y的切线设为0.0(与原输出类型匹配的零值)
    tangents_out = (t0 + t1, 0.0)

    return primals_out, tangents_out

if __name__ == "__main__":
    x = 1.
    y = 2.
    print(jacobian(func)(x, y))

方法二:在原函数中使用stop_gradient

如果某个输出从一开始就不需要微分,可以直接在原函数中用jax.lax.stop_gradient包装,无需自定义JVP:

from jax import jacobian, lax

def func(x, y):
    return x+y, lax.stop_gradient(x*y)

if __name__ == "__main__":
    x = 1.
    y = 2.
    print(jacobian(func)(x, y))

两种方法运行后都会输出:

(DeviceArray(1., dtype=float32, weak_type=True), DeviceArray(0., dtype=float32, weak_type=True))

对应第一个输出x+y的雅可比矩阵(对x、y的导数均为1),第二个输出因被排除在微分计算外,雅可比矩阵为0。

内容的提问来源于stack exchange,提问作者Jingyang Wang

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 02:33:15