Python列表推导式版本计算平均球距离更慢?求优化方案
平均球距离计算的性能优化问题(针对大规模稀疏矩阵)
计算逻辑
计算"平均球距离"的步骤如下:
- 选择一个节点
- 选取该节点距离为d的所有元素(距离d的shell)
- 找到距离d的所有节点周围半径为d的shell
- 计算这些距离d、半径d的shell的平均距离
当前实现使用稀疏矩阵存储节点邻接关系,借助scipy的Dijkstra算法计算最短路径;由于仅使用矩阵子集元素,通过字典映射元素与最短路径列表的索引。
当前实现代码
导入依赖库
import numpy as np import time import random from random import randint from scipy.sparse import csr_matrix from scipy.sparse.csgraph import breadth_first_tree import scipy as scp from scipy.stats import uniform from scipy import io from scipy import sparse # 稀疏矩阵工具包 from scipy import linalg # 线性代数工具包 from scipy.sparse import identity from scipy.sparse import csr_matrix from scipy.sparse.csgraph import breadth_first_order from scipy.sparse import csr_matrix from scipy.sparse.csgraph import shortest_path,dijkstra
读取稀疏矩阵
input_m = "L-4-4.0-0.02-4-2-1.mtx" L = scp.sparse.csc_matrix(scp.io.mmread(input_m), dtype=int) ID = identity(np.shape(L)[0], dtype='int8', format='dia') WA = abs(5*ID - L)
定义Shell/Ball获取函数
def GetShellBall(WA,n,ind): p0 = np.zeros(np.shape(L)[0]) p0[ind] = 1 newp = p0 ball = [] shell = [] ball.append(ind) shell.append([ind]) for it in range(n): newp = WA@newp for it2 in np.where(newp)[0]: if it2 in ball: newp[it2] = 0 else: ball.append(it2) shell.append(np.where(newp)[0]) return ball,shell def GetShellj(WA,n,ind): p0 = np.zeros(np.shape(L)[0]) p0[ind] = 1 newp = p0 ball = [] shell = [] ball.append(ind) shell.append([ind]) for it in range(n): newp = WA@newp for it2 in np.where(newp)[0]: if it2 in ball: newp[it2] = 0 else: ball.append(it2) shell.append(np.where(newp)[0]) return shell[n]
生成Shell/Ball与最短路径
%%time it = 0 N = 11 ball,shell = GetShellBall(WA,N,it) # 节点it周围距离N内的ball与shell dict_ba = dict(zip(ball, np.arange(len(ball)))) # 映射:ball元素 -> 列表索引 dict_ab = dict(zip(np.arange(len(ball)), ball)) # 映射:列表索引 -> ball元素 spaths = np.asarray([dijkstra(csgraph=WA, directed=True, limit = N,unweighted = True, indices=it, return_predecessors=False) for it in ball])
CPU耗时:用户态7.61秒,系统态11.9毫秒,总计7.62秒; wall time:7.63秒
核心计算性能测试
此前列表推导式版本的性能问题已通过修复bug解决,当前测试结果如下:
循环版本
%%timeit #for i in range(len(shell)): DD = [] for i in range(int(np.floor(N/2))): sumd = [] chosen_paths = spaths[[dict_ba[it] for it in shell[i]]] for eli in shell[i]: shellj = GetShellj(WA,i,eli) sumd = [[chp[elj] for elj in shellj] for chp in chosen_paths] DD.append(np.mean(sumd))
每次循环耗时36.2 ms ± 943 µs(7次运行、每次10个循环的均值±标准差)
列表推导式版本
%%timeit #for i in range(len(shell)): DD = [np.mean([ [ [ chp[elj] for elj in GetShellj(WA,i,eli) ] for chp in spaths[[dict_ba[it] for it in shell[i]]] ] for eli in shell[i]]) for i in range(int(np.floor(N/2)))]
每次循环耗时440 ms ± 8.05 ms(7次运行、每次1个循环的均值±标准差)
优化需求
寻求针对**大规模稀疏矩阵(约50万×50万)**场景的更高效实现方案。
内容的提问来源于stack exchange,提问作者Kregnach
相关产品推荐
相关产品推荐

