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的核心是指定映射轴,默认会将输入的第一个轴作为拆分维度,把输入拆成多个子数组逐个传入被包装函数,再堆叠结果。你的错误尝试存在两个核心问题:
- 跳过关键计算步骤:代码中定义了
func1却未调用,直接对原始data执行两次嵌套的求和vmap,完全跳过了原函数中必须的cumprod步骤,数值自然错误。 - 映射轴与计算逻辑不匹配:
func2 = vmap(jnp.sum)会把输入的第一个轴拆分,对每个子数组做全局求和。比如输入(31,7)的数组,会拆成31个(7,)的行,每行求和得到单个数值,最终返回(31,)的结果。func3 = vmap(func2)会把data的第一个轴(11个(31,7)的数据集)拆分,每个数据集经func2处理后返回(31,),最终堆叠成(11,31)的结果——若你说维度符合预期,大概率是误写了vmap的轴参数,但核心逻辑仍不符合原函数要求。
正确做法需要分步嵌套的原因
原函数的逻辑是三层递进的操作,分步嵌套vmap完全对应了这个逻辑:
- 外层vmap:
vmap(inner)(data)对应原函数的外层循环,遍历data的第一个轴(11个数据集),对每个数据集执行完整的处理逻辑。 - 内层cumprod的vmap:
vmap(jnp.cumprod)(dataset)对应原函数中对单数据集的每行做累积乘积,把dataset的第一个轴(31行)作为映射轴,每行独立执行cumprod,得到和原函数一致的(31,7)结果。 - 求和步骤:
jnp.sum(cp, axis=1)直接对应原函数中对行维度求和的要求,得到(7,)的结果,最终11个结果堆叠成(11,7)的输出。
分步嵌套的关键在于:每一层vmap都明确对应原循环的一层,确保每个操作的维度、顺序和原函数完全对齐。而你的错误尝试直接嵌套vmap,既跳过了核心计算步骤,又让映射轴偏离了原函数的逻辑,最终导致数值错误。
内容的提问来源于stack exchange,提问作者hasco641
相关产品推荐
相关产品推荐

