如何在JAX中对不等长数组列表实现函数向量化
问题描述
给定以下JAX函数:
import jax.numpy as jnp def test(x): return jnp.sum(x)
尝试用jax.vmap向量化该函数:
v_test = jax.vmap(test)
输入为一组不等长数组:
x1 = jnp.array([1,2,3]) x2 = jnp.array([4,5,6,7]) x3 = jnp.array([8,9]) x4 = jnp.array([10]) x = [x1, x2, x3, x4]
执行v_test(x)时触发错误:
ValueError: vmap got inconsistent sizes for array axes to be mapped: the tree of axis sizes is: ([3, 4, 2, 1],)
需要在不填充数组的前提下,实现对不等长数组列表批量应用test函数。
解决方案
方法1:使用jax.lax.map
jax.lax.map专门适配结构一致但元素维度/长度可变的输入,会遍历序列中每个元素并应用目标函数,无需强制所有元素长度统一。代码示例:
import jax import jax.numpy as jnp def test(x): return jnp.sum(x) x1 = jnp.array([1,2,3]) x2 = jnp.array([4,5,6,7]) x3 = jnp.array([8,9]) x4 = jnp.array([10]) x = [x1, x2, x3, x4] result = jax.lax.map(test, x) print(result) # 输出: [ 6 22 17 10]
方法2:列表推导配合JAX自动微分
如果不需要批量操作的底层优化,直接用Python列表推导也能实现,JAX会自动追踪其中的计算过程:
result = jnp.array([test(arr) for arr in x]) print(result) # 输出: [ 6 22 17 10]
为什么vmap无法解决?
jax.vmap的设计目标是处理形状规则的批量输入,要求被映射的轴在所有输入元素中尺寸一致,而不等长数组不满足该条件,因此会触发尺寸不匹配的错误。
内容的提问来源于stack exchange,提问作者MOON
相关产品推荐
相关产品推荐

