JAX中vmap处理多参数的异常行为及报错原因咨询
JAX vmap多参数映射时尺寸不匹配错误的原因解析
报错核心原因
你写的f2用@partial(vmap, in_axes=(None, 0, 0))时,vmap的逻辑是对y和z的第0轴做同批次并行迭代——也就是每次迭代会同时取y[i]和z[i]代入函数计算,这就要求y和z的第0轴长度必须完全一致。但你的y是长度5的数组,z是长度3的数组,两个映射轴尺寸不匹配,因此触发ValueError。
为什么f1能正常运行?
f1是通过手动添加None扩展维度,利用JAX的广播机制实现笛卡尔积式的计算:
x[:, None, None]把x从(7,)扩展为(7,1,1)y[None, :, None]把y从(5,)扩展为(1,5,1)z[None, None, :]把z从(3,)扩展为(1,1,3)
三个数组通过广播自动对齐为(7,5,3)的形状,本质是遍历x的7个元素、y的5个元素、z的3个元素的所有组合进行计算,和vmap的“同批次迭代”逻辑完全不同。
用vmap实现f1效果的正确方式
如果想用vmap复刻f1的结果,需要用嵌套vmap分别处理y和z的不同维度:
from functools import partial import jax import jax.numpy as jnp @partial(jax.vmap, in_axes=(None, None, 0), out_axes=2) # 映射z的第0轴,输出对应维度2 @partial(jax.vmap, in_axes=(None, 0, None), out_axes=1) # 映射y的第0轴,输出对应维度1 def f2(x, y, z): return x * z + y x = jnp.arange(7) y = jnp.arange(5) z = jnp.arange(3) print(f2(x, y, z).shape) # 输出(7, 5, 3),和f1一致
这里内层vmap遍历y的每个元素,外层vmap遍历z的每个元素,最终实现三个维度的笛卡尔积计算。
内容的提问来源于stack exchange,提问作者Jingyang Wang
相关产品推荐
相关产品推荐

