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

