PyTorch/Numpy中小向量与超大稀疏矩阵乘法的优化咨询
针对小向量a与超大稀疏矩阵b的乘法优化需求(b只读、需多次和不同a计算),以下是基于稀疏特性的实用优化方案,附代码示例:
方案1:转换为CSC稀疏矩阵(开发效率最高)
因为a @ b本质是对b的每一列做加权求和,而CSC(压缩稀疏列)格式天然适配列操作。scipy.sparse的CSC矩阵会自动跳过零元素,只计算非零项的贡献,预处理仅需一次,后续计算速度远快于密集矩阵乘法。
import numpy as np from scipy.sparse import csc_matrix # 从memmap加载原始b矩阵(假设已提前保存为memmap文件) B = 32 M = 10000000 b_memmap = np.memmap('b_matrix.npy', dtype=bool, mode='r', shape=(B, M)) # 预处理:转换为CSC格式(仅执行一次) b_csc = csc_matrix(b_memmap) # 多次执行与不同a的乘法 for _ in range(100): a = np.random.rand(B) result = a @ b_csc # 自动利用稀疏性优化计算 # 处理结果...
方案2:自定义稀疏存储(针对极小B的轻量优化)
当B很小(比如你的场景是32),可以手动提取每一列的非零索引,用更紧凑的格式存储,再通过Numba JIT加速循环计算,比scipy.sparse的开销更低。
import numpy as np from numba import njit B = 32 M = 10000000 # 从memmap加载b b_memmap = np.memmap('b_matrix.npy', dtype=bool, mode='r', shape=(B, M)) # 预处理:提取非零索引的紧凑存储(仅执行一次) # 用indptr记录每一列非零元素的起始位置,indices存储所有非零元素的行索引 counts = np.array([np.sum(b_memmap[:, j]) for j in range(M)], dtype=np.int32) indptr = np.zeros(M + 1, dtype=np.int64) indptr[1:] = np.cumsum(counts) indices = np.zeros(indptr[-1], dtype=np.int32) pos = 0 for j in range(M): cols = np.where(b_memmap[:, j])[0] indices[pos:pos + len(cols)] = cols pos += len(cols) # Numba JIT并行加速计算 @njit(parallel=True) def compute_result(a, indices, indptr, M): result = np.zeros(M, dtype=np.float64) for j in range(M): start = indptr[j] end = indptr[j + 1] total = 0.0 for k in range(start, end): total += a[indices[k]] result[j] = total return result # 多次计算 for _ in range(100): a = np.random.rand(B) result = compute_result(a, indices, indptr, M) # 处理结果...
方案3:位掩码+向量指令(极致性能优化)
因为B=32刚好匹配32位整数的位数,可以把每一列的bool值打包成一个32位整数(位掩码),内存占用极低(仅40MB),再通过Numba编译出利用CPU向量指令的代码,速度达到最优。
import numpy as np from numba import njit, uint32, float64 B = 32 M = 10000000 # 从memmap加载b b_memmap = np.memmap('b_matrix.npy', dtype=bool, mode='r', shape=(B, M)) # 预处理:将每一列转换为32位位掩码(仅执行一次) bitmasks = np.zeros(M, dtype=np.uint32) for j in range(M): mask = 0 for i in range(B): if b_memmap[i, j]: mask |= (1 << i) bitmasks[j] = mask # Numba JIT并行加速,利用位运算快速求和 @njit(parallel=True) def compute_with_bitmask(a, bitmasks, B, M): result = np.zeros(M, dtype=np.float64) for j in range(M): mask = bitmasks[j] total = 0.0 i = 0 while mask: if mask & 1: total += a[i] mask >>= 1 i += 1 result[j] = total return result # 多次计算 for _ in range(100): a = np.random.rand(B) result = compute_with_bitmask(a, bitmasks, B, M) # 处理结果...
方案选择建议
- 追求开发效率:直接用方案1,代码最简,无需手动处理稀疏逻辑。
- 追求极致性能:选方案3,内存占用最少,速度最快,完美适配B=32的场景;方案2适合B稍大的情况。
- 所有预处理步骤仅需执行一次,后续和不同a的计算都能复用预处理后的结构,大幅降低重复计算量。
内容的提问来源于stack exchange,提问作者Garvey
相关产品推荐
相关产品推荐

