JAX中带custom_vjp的vmapped函数调用vjp失败的问题咨询
问题分析与解决
你遇到的问题是在结合custom_vjp和vmap时,JAX抛出了形状不匹配的错误,且错误提示存在误导性。实际错误根源在于自定义VJP的前向函数中递归调用了被装饰后的函数,导致JAX无法正确跟踪批量维度下的形状关系。
错误原因
test_func_fwd中调用了test_func(f, primal),而test_func已经被custom_vjp装饰,这会触发无限递归,同时干扰JAX对vmap后函数输出/输入形状的正确推导,最终导致错误提示混淆了余切(应匹配输出形状)与原始输入的形状要求。
解决方案
修改前向函数,直接计算输出值而非调用被装饰后的函数,避免递归并让JAX正确跟踪形状。
修正后的代码
from functools import partial import jax import jax.numpy as jnp from jax import custom_vjp, vjp, vmap from jax._src.typing import Array, Callable @partial(custom_vjp, nondiff_argnums=(0,)) def test_func(f: Callable[..., float], R: Array ) -> float: return f(jnp.dot(R, R)) def test_func_fwd(f, primal): # 直接计算输出,避免递归调用被custom_vjp装饰后的test_func primal_out = f(jnp.dot(primal, primal)) residual = 2. * primal * primal_out return primal_out, residual def test_func_bwd(f, residual, cotangent): cotangent_out = residual * cotangent return (cotangent_out, ) test_func.defvjp(test_func_fwd, test_func_bwd) test_func = vmap(test_func, in_axes=(None, 0)) if __name__ == "__main__": def f(x): return x # vjp primal, f_vjp = vjp(partial(test_func, f), jnp.ones((10, 3)) ) cotangent = jnp.ones(10) cotangent_out = f_vjp(cotangent) print(cotangent_out[0].shape) # 输出 (10, 3),符合预期
验证结果
修正后代码运行会输出(10, 3),这是正确的输入余切形状:输入为(10,3),输出为(10,),拉回函数接受(10,)的余切,返回与输入同形状的(10,3)余切。
内容的提问来源于stack exchange,提问作者Jingyang Wang
相关产品推荐
相关产品推荐

