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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 18:10:36