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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 03:34:59