Python(NumPy):使用花式索引实现内存高效的数组乘法
解决方案:避免中间数组的索引式矩阵乘法
你的核心问题是避免生成A[I]这个超大中间数组,同时完成每个A[I[k]]与B[k]的矩阵乘法。以下是几种高效的NumPy/Numba实现方案:
方案1:分组批量计算(纯NumPy)
利用np.unique提取唯一索引,将B按索引分组后分别与对应的A子矩阵相乘,最后合并结果。这种方式不会生成完整的A[I],内存占用仅为分组后的B块总和。
import numpy as np from numpy.random import rand, randint A = rand(1000,5,5) B = rand(40000000,5,1) I = randint(low=0, high=1000, size=40000000) unique_I, inv_I = np.unique(I, return_inverse=True) result = np.empty_like(B) # 遍历每个唯一索引,批量计算对应组的乘法 for idx in unique_I: mask = I == idx result[mask] = A[idx] @ B[mask]
优缺点:无需额外依赖,内存占用低;若唯一索引数量接近1000,循环次数较多,速度中等。
方案2:Numba并行加速(最高效)
用Numba将循环编译为机器码,直接逐个计算每个A[I[k]] @ B[k],完全避免中间数组,内存占用最小,且并行模式下速度远超纯NumPy方案。
import numpy as np from numpy.random import rand, randint from numba import njit, prange A = rand(1000,5,5) B = rand(40000000,5,1) I = randint(low=0, high=1000, size=40000000) @njit(parallel=True) def indexed_matmul(A, B, I, out): total = len(I) # 并行遍历每个元素 for k in prange(total): idx = I[k] # 手动计算5x5矩阵与5x1向量的乘法 for row in range(5): val = 0.0 for col in range(5): val += A[idx, row, col] * B[k, col, 0] out[k, row, 0] = val result = np.empty_like(B) indexed_matmul(A, B, I, result)
优缺点:速度最快,内存占用极低;需要安装Numba库,代码需写显式循环(但对于固定的5x5矩阵,循环代码量很小)。
方案3:利用逆索引重组结果
如果分组计算后需要保持原顺序,可通过inv_I直接重组结果块,避免多次创建mask:
import numpy as np from numpy.random import rand, randint A = rand(1000,5,5) B = rand(40000000,5,1) I = randint(low=0, high=1000, size=40000000) unique_I, inv_I = np.unique(I, return_inverse=True) result_blocks = [] for i in range(len(unique_I)): # 提取所有对应当前唯一索引的B元素 block_B = B[inv_I == i] # 计算矩阵乘法 result_blocks.append(A[unique_I[i]] @ block_B) # 拼接所有结果块,自动恢复原顺序 result = np.concatenate(result_blocks)
优缺点:与方案1逻辑类似,但减少了mask的重复创建,代码更简洁;性能与方案1接近。
为什么原方案内存溢出?
A[I]会生成形状为(40000000,5,5)的数组,按float64计算,总内存为40000000 * 5 * 5 * 8 = 8GB,远超普通机器的内存容量,因此必须避免生成该中间数组。
内容的提问来源于stack exchange,提问作者grescha
相关产品推荐
相关产品推荐

