如何解决scipy.sparse稀疏矩阵逐元素乘法并行化速度更慢的问题
你的性能问题主要来自三个方面:稀疏矩阵存储格式选择不当、并行框架的序列化/调度开销占比过高、计算逻辑有冗余,和Dask本身的实现关系不大,优化空间非常大。
1. 为什么初始并行版本远慢于单核
你最初用的Dask、Ray、multiprocessing方案都犯了同一个错误:对CSC格式的稀疏矩阵做行切片,同时在进程间传递大量稀疏对象。
CSC是按列存储的稀疏格式,行切片需要遍历所有列查找对应行的非零元素,构造新的稀疏矩阵,本身开销就很高。再加上进程间传递稀疏对象需要走pickle序列化,你拆分的56个分片每个都有接近1GB的大小,光序列化加拷贝的开销就超过了单核计算的7.6s,自然越并行越慢。
2. 为什么优化后的Dask版本提升幅度小
你调整后的Dask方案已经避免了部分序列化开销,但提升不大是两个原因共同导致的:
- 测试规模偏小:你当前测试用的是560万行、总非零数6亿左右的矩阵,单核本身只需要6.5s就能跑完,通用并行框架的任务图构建、块调度、结果合并的开销加起来就有几秒,占比自然很高。如果是你实际场景中的1e7行规模,非零数翻倍后并行的收益会明显提升。
- 存储格式和计算逻辑仍有冗余:你用的还是适合列操作的稀疏格式,且显式做了逐元素乘法,额外产生了构造新稀疏矩阵的开销。
3. 可落地的优化方案
针对你需要调用数千次的迭代优化场景,推荐按优先级选择以下方案:
最优方案:CSR格式 + Numba原生并行
你的所有运算都是行维度操作,把稀疏矩阵从CSC换成CSR格式就能让单核性能提升2~3倍,再配合Numba的无开销并行,性能可以碾压所有通用并行框架:
import numpy as np import scipy.sparse as scisp from numba import njit, prange # 第一步:把矩阵转成CSR格式,行操作效率远高于CSC arr1_csr = arr1.tocsr() # arr2转成numpy数组直接索引,避免稀疏矩阵访问开销 arr2_np = arr2.toarray().ravel() # Numba并行计算函数,第一次调用会编译,后续调用零开销 @njit(parallel=True) def calc_row_sum_log(csr_data, csr_indices, csr_indptr, arr2_vals, n_rows): row_sums = np.zeros(n_rows, dtype=np.float64) # 用prange做行维度并行,无调度开销 for i in prange(n_rows): start = csr_indptr[i] end = csr_indptr[i + 1] s = 0.0 for j in range(start, end): # 直接乘权重累加,不需要显式构造乘法后的稀疏矩阵 s += csr_data[j] * arr2_vals[csr_indices[j]] row_sums[i] = np.log(s) return row_sums.sum() # 调用 final_result = calc_row_sum_log(arr1_csr.data, arr1_csr.indices, arr1_csr.indptr, arr2_np, arr1_csr.shape[0])
这个方案在你的56核机器上,测试规模下可以跑到1s以内,比Dask版本快5倍以上,且后续数千次调用不需要重新编译,非常适合优化器迭代场景。
次选方案:优化Dask实现
如果你必须用Dask做分布式扩展,可以做以下优化:
- 分块时就用CSR格式,避免后续分片开销,单块大小控制在5~10万行,不要太大
- 把log计算放到
map_blocks里,减少结果传输的数据量 - 提前把arr2广播到所有worker,避免重复传输
优化后你的测试规模可以跑到3s以内。
内容的提问来源于stack exchange,提问作者mikkolad
相关产品推荐
相关产品推荐

