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

JAX中嵌套vmap映射与分步vmap应用的差异及问题解析

问题分析:JAX中vmap嵌套调用的错误逻辑与正确解法

背景与目标

我们需要用JAX的vmap替代循环,复现以下numpy风格的函数:

data = random.normal(key, shape=(11, 31, 7))

def ex2_numpy_equivalent(data):
    result = []
    for d in data: 
        cp = jnp.cumprod(d, axis=-1)
        s = jnp.sum(cp, axis=1)
        result.append(s)
    return jnp.stack(result)

Eric给出的正确解法:

def loopless_loop_ex2(data):
    """Data is three-dimensional of shape (n_datasets, n_rows, n_columns)"""

    def inner(dataset):
        """dataset is two-dimensional of shape (n_rows, n_columns)"""
        cp = vmap(jnp.cumprod)(dataset)
        s = jnp.sum(cp, axis=1)  # 注:原示例中vmap(jnp.sum)应为笔误,实际需对应原函数的axis=1求和
        return s 

    return vmap(inner)(data)

我的错误尝试:

func1 = vmap(jnp.cumprod)
func2 = vmap(jnp.sum)
func3 = vmap(func2)

func3(data)

该尝试输出维度符合预期,但数值完全错误。下面分析错误逻辑与正确做法的核心原因。

错误尝试的运行逻辑

vmap的核心是指定映射轴,默认会将输入的第一个轴作为拆分维度,把输入拆成多个子数组逐个传入被包装函数,再堆叠结果。你的错误尝试存在两个核心问题:

  1. 跳过关键计算步骤:代码中定义了func1却未调用,直接对原始data执行两次嵌套的求和vmap,完全跳过了原函数中必须的cumprod步骤,数值自然错误。
  2. 映射轴与计算逻辑不匹配:
    • func2 = vmap(jnp.sum)会把输入的第一个轴拆分,对每个子数组做全局求和。比如输入(31,7)的数组,会拆成31个(7,)的行,每行求和得到单个数值,最终返回(31,)的结果。
    • func3 = vmap(func2)会把data的第一个轴(11个(31,7)的数据集)拆分,每个数据集经func2处理后返回(31,),最终堆叠成(11,31)的结果——若你说维度符合预期,大概率是误写了vmap的轴参数,但核心逻辑仍不符合原函数要求。

正确做法需要分步嵌套的原因

原函数的逻辑是三层递进的操作,分步嵌套vmap完全对应了这个逻辑:

  1. 外层vmap:vmap(inner)(data)对应原函数的外层循环,遍历data的第一个轴(11个数据集),对每个数据集执行完整的处理逻辑。
  2. 内层cumprod的vmap:vmap(jnp.cumprod)(dataset)对应原函数中对单数据集的每行做累积乘积,把dataset的第一个轴(31行)作为映射轴,每行独立执行cumprod,得到和原函数一致的(31,7)结果。
  3. 求和步骤:jnp.sum(cp, axis=1)直接对应原函数中对行维度求和的要求,得到(7,)的结果,最终11个结果堆叠成(11,7)的输出。

分步嵌套的关键在于:每一层vmap都明确对应原循环的一层,确保每个操作的维度、顺序和原函数完全对齐。而你的错误尝试直接嵌套vmap,既跳过了核心计算步骤,又让映射轴偏离了原函数的逻辑,最终导致数值错误。

内容的提问来源于stack exchange,提问作者hasco641

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 01:41:04