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

JAX中sum生成大中间数组拖慢GPU性能的优化方案问询

优化批量列索引求和:避免大中间数组

你的核心问题是批量处理列索引求和时,vmap会生成B×C的中间数组,导致GPU内存瓶颈。下面提供两种高效重构方式,无需生成大中间数组:

方法1:用lax.fori_loop按列循环累加

这种方式逐列处理数据,每次仅生成B大小的临时数组(单列的批量取值),最终累加得到每个batch的总和,内存复杂度为O(B)而非O(B×C),适合大规模场景。

代码实现:

import time
import numpy as np
import jax
import jax.numpy as jnp
from jax import lax

batch_size = 10_000
data = jnp.array(np.random.random((400, 10000)))
batch_to_sum = jnp.array(np.random.randint(data.shape[0], size=(batch_size, data.shape[1])))

def batch_sum_fori(data, batch_to_sum):
    B, C = batch_to_sum.shape
    def loop_body(col_idx, accumulator):
        # 提取当前列的数据,按批量索引取值后累加
        col_vals = data[:, col_idx]
        batch_col_vals = col_vals[batch_to_sum[:, col_idx]]
        return accumulator + batch_col_vals
    # 初始化累加器为全0,循环遍历所有列
    return lax.fori_loop(0, C, loop_body, init_val=jnp.zeros(B))

func_fori = jax.jit(batch_sum_fori)

print(" *** fori_loop ***")
t0 = time.time()
res_fori = func_fori(data, batch_to_sum).block_until_ready()
print(f"fori_loop + jit 1st pass: {(time.time() - t0)}")
t0 = time.time()
res_fori = func_fori(data, batch_to_sum).block_until_ready()
print(f"fori_loop + jit 2nd pass: {(time.time() - t0)}")

# 验证结果一致性
assert jnp.allclose(res, res_fori)

方法2:利用高级索引的隐式优化(XLA自动融合)

如果你不想写循环,可以直接调整索引方式,让XLA自动融合求和与gather操作,避免显式存储中间数组。虽然代码看起来和原方法类似,但JAX的XLA编译器可能会优化掉中间数组的生成(取决于场景):

def batch_sum_direct(data, batch_to_sum):
    C = data.shape[1]
    # 直接索引后沿列维度求和,XLA可能会融合操作
    return jnp.sum(data[batch_to_sum, jnp.arange(C)], axis=1)

func_direct = jax.jit(batch_sum_direct)

不过这种方式的优化效果依赖于XLA的自动融合能力,当B×C规模极大时,fori_loop的内存控制更可靠。

效果对比

对于batch_size=1e7、C=1e4的场景,fori_loop方法的内存占用仅为原vmap方法的1/C,能有效避免GPU内存瓶颈,同时保持相近的计算速度(jit编译后)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 08:13:19