JAX中vjp调用含custom_vjp的cart_deriv时提示缺少cotangent参数
JAX custom_vjp 反向传播参数错误问题排查与修复
问题描述
我编写了一个JAX函数cart_deriv(),用于计算输入函数f的笛卡尔导数,代码如下:
from functools import partial import jax from jax import custom_vjp, jacrev, vjp, jnp from jax.tree_util import Partial from typing import Callable, Array @partial(custom_vjp, nondiff_argnums=0) def cart_deriv(f: Callable[..., float], l: int, R: Array ) -> Array: df = lambda R: f(l, jnp.dot(R, R)) for i in range(l): df = jacrev(df) return df(R) def cart_deriv_fwd(f, l, primal): primal_out = cart_deriv(f, l, primal) residual = cart_deriv(f, l+1, primal) ## 测试用残差 return primal_out, residual def cart_deriv_bwd(f, residual, cotangent): cotangent_out = jnp.ones(3) ## 测试用输出 return (None, cotangent_out) cart_deriv.defvjp(cart_deriv_fwd, cart_deriv_bwd) if __name__ == "__main__": def test_func(l, r2): return l + r2 primal_out, f_vjp = vjp(cart_deriv, Partial(test_func), 2, jnp.array([1., 2., 3.]) ) cotangent = jnp.ones((3, 3)) cotangent_out = f_vjp(cotangent) print(cotangent_out[1].shape)
运行时触发错误:
TypeError: cart_deriv_bwd() missing 1 required positional argument: 'cotangent'
错误原因
问题出在自定义VJP反向函数的签名不符合JAX规范:
- 使用
partial(custom_vjp, nondiff_argnums=0)时,JAX会将所有非微分参数打包成一个元组,作为反向函数的第一个参数传入,而非单个参数。 - 你的
cart_deriv_bwd把第一个参数写成了单个f,导致参数匹配错位:原本的residual被当作cotangent,自然缺失了最后一个参数。
修复方案
调整反向函数的签名,让第一个参数接收非微分参数的元组,再从中取出f:
def cart_deriv_bwd(nondiff_args, residual, cotangent): # 从元组中取出非微分参数f f = nondiff_args[0] # 这里可根据实际需求计算l和R的余切,测试用值保持不变 cotangent_l = None # l为整数,若无需微分可返回None cotangent_R = jnp.ones(3) return (cotangent_l, cotangent_R)
修改后,JAX能正确匹配反向函数的参数,错误即可解决。
内容的提问来源于stack exchange,提问作者Jingyang Wang
相关产品推荐
相关产品推荐

