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

Jax vmap中in_axes参数搭配关键字参数调用报错的解决方法问询

Jax vmap中in_axes参数搭配关键字参数调用报错的解决方法问询

我完全懂你遇到的这个头疼问题!确实JAX的vmap在混用元组形式的in_axes和关键字参数时会触发无消息的AssertionError,因为元组是严格对应位置参数顺序的,一旦用关键字传参,内部参数结构解析就会出现不匹配。不过有两种实用的解决办法:

方法一:使用字典形式的in_axes指定参数轴映射

这是最直接优雅的方案,JAX支持用字典来明确绑定每个参数名和对应的轴信息,不管参数是位置传递还是关键字传递都能正确识别。修改后的代码如下:

from jax import vmap
import numpy as np

def foo(a, b, c):
    return a * b + c

# 用字典指定每个参数的in_axes,键是参数名,值是对应轴
foo_vmap = vmap(foo, in_axes={'a': 0, 'b': 0, 'c': None})

aj, bj = np.random.rand(2, 100, 1)
foo_vmap(aj, bj, c=10)  # 关键字参数调用正常运行
foo_vmap(aj, bj, 10)    # 位置参数调用也依然生效

方法二:包装函数统一参数传递形式

如果不想修改in_axes的格式,可以写一个包装函数,把关键字参数转为位置参数的形式传递给原函数,这样vmap的元组in_axes就能正常工作了:

from jax import vmap
import numpy as np

def foo(a, b, c):
    return a * b + c

# 包装函数,接收关键字参数并转为位置参数传递
def foo_wrapper(a, b, c=10):
    return foo(a, b, c)

foo_vmap = vmap(foo_wrapper, in_axes=(0, 0, None))

aj, bj = np.random.rand(2, 100, 1)
foo_vmap(aj, bj, c=10)  # 现在可以正常执行

为什么原来的代码会报错?

当你用元组(0, 0, None)作为in_axes时,vmap默认按位置参数的顺序来匹配轴信息。但当你用c=10这种关键字参数传递时,JAX内部在解析参数树结构时,会出现参数数量和轴映射数量不匹配的情况,最终触发了那个无消息的AssertionError。而字典形式的in_axes直接绑定参数名和轴,就能避开这个问题。

备注:内容来源于stack exchange,提问作者Amith M

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.23 08:39:07