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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.31 06:09:10