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
相关产品推荐
相关产品推荐

