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

Jax生成指定随机矩阵时内存溢出问题排查与优化咨询

内存占用过高的原因
  • 你的实现中,jax.random.choice(..., replace=False)的无放回采样默认采用Fisher-Yates洗牌算法,该算法需要为每行生成一个大小为N的掩码/临时数组跟踪已选元素。当N=50000、N行并行处理时,临时内存总占用量达到N*N=2.5e9个元素,直接触发内存溢出。
  • jax.vmap对每行独立调用采样函数,未利用JAX的向量化优化共享临时内存,反而累积了每行的O(N)临时空间开销。
正确实现方式

针对每行选k个不同元素(k远小于N)的场景,我们可以实现仅做k步的Fisher-Yates采样,避免生成整个N长度的排列或掩码,将每行的临时内存开销从O(N)降到O(k):

import jax
import jax.numpy as jnp

def sample_k_unique(N, k, key):
    # 初始化数组存储前k个待选元素
    arr = jnp.arange(k)
    # 生成k个子密钥用于每一步采样
    step_keys = jax.random.split(key, k)
    
    # 执行k步Fisher-Yates洗牌,仅生成前k个唯一元素
    for i in range(k):
        # 从i到N-1中随机选一个位置j
        j = jax.random.randint(step_keys[i], shape=(), minval=i, maxval=N)
        # 交换当前位置i与j的元素:若j >=k,arr中无j,直接用j替换;否则交换arr[i]和arr[j]
        arr = arr.at[i].set(j if j >= k else arr[j])
        if j < k:
            arr = arr.at[j].set(i)
    return arr

# 向量化批量生成每行
batch_sample = jax.vmap(sample_k_unique, in_axes=(None, None, 0))

def generate_mat(N, k, key=jax.random.PRNGKey(0)):
    # 生成N个独立密钥,对应每行的采样
    row_keys = jax.random.split(key, N)
    return batch_sample(N, k, row_keys)

优化原理

  • 仅执行k步Fisher-Yates洗牌,无需生成完整的N长度排列,每行临时内存仅为O(k)级别。
  • 利用JAX的vmap实现批量处理,同时通过局部交换逻辑避免了O(N)掩码数组的生成。

如果k与N接近(比如k>N/2),可以反过来采样N-k个元素,然后取补集,进一步节省计算资源。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 05:43:22