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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 12:35:19