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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 18:40:13