如何对zip打包的参数使用Jax vmap实现向量化映射?
错误原因
jax.vmap 和Python内置map的入参规则完全不同:
- 内置
map会遍历传入的可迭代对象,逐个把元素传给目标函数执行 vmap要求传入的是沿第0轴堆叠的数组(或数组构成的PyTree结构),它会自动沿第0轴切分数组传入函数做批计算,不会自动解析Python原生的zip迭代器、普通列表这类结构。
你直接传入zip(xs, ys)返回的Python迭代器时,vmap无法识别其中的批处理维度,就会抛出参数秩为0的报错。
正确写法
首先把分散的数组合并为带批维度的堆叠数组,再传入vmap即可。推荐直接把函数改成接收两个独立参数的形式,vmap会自动对所有入参的第0轴做同步映射,写法最简洁:
import jax import jax.numpy as jnp def f(x, y): return x.sum() + y.sum() # 直接构造带批维度的数组,批大小为4 xs = jnp.zeros((4, 3)) ys = jnp.zeros((4, 2)) jax.vmap(f)(xs, ys)
运行输出和原生map的计算结果完全一致:
DeviceArray([0., 0., 0., 0.], dtype=float32)
如果你要保留原函数接收单个元组参数的写法,也可以先把列表堆叠为数组后打包成元组传入,不要用原生zip:
def f(x_y): x, y = x_y return x.sum() + y.sum() # 把原列表堆叠为带批维度的数组 xs_stacked = jnp.stack([jnp.zeros(3) for i in range(4)]) ys_stacked = jnp.stack([jnp.zeros(2) for i in range(4)]) jax.vmap(f)((xs_stacked, ys_stacked))
补充说明:vmap返回的是沿第0轴堆叠好的单个数组,原生map返回的是独立结果组成的列表,二者数值等价。如果需要和原示例完全一致的列表输出,对vmap返回结果做一次遍历转换即可。
内容的提问来源于stack exchange,提问作者marius
相关产品推荐
相关产品推荐

