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

如何优化Python中1000×1000稀疏矩阵乘法运算以提升速度

1000×1000规模矩阵运算代码优化方案

根据梯度方程,矩阵乘法需同时使用@(矩阵乘法)和*(逐元素乘法)运算。以下是待优化的原代码:

# parameters
beta     =  0.98 
alpha    =  0.03
delta    =  0.1
T        =  1000
loop     =  1
dif      =  1
tol      =  1e-8

kss      =  ((1 / beta - (1 - delta)) / alpha)**(1 / (alpha - 1))
k        =  np.linspace(0.5 * kss, 1.8 * kss, T)

k_reshaped   =  k.reshape(-1, 1)
c            =  k_reshaped ** alpha + (1 - delta) * k_reshaped - k
c[c<0]       =  1e-11
c            =  np.log(c)
beta_square  =  beta**2

# multiplication
I   =  np.identity(T)
E   =  np.ones(T)[:,None]
Q2  =  I

while np.any(dif > tol) and loop < 200:
    J   =  beta * Q2
    B   =  inv(I - J)

    Q3  =  np.zeros([T,T])
    ini =  np.argmax(c + (B @ (J * c) @ E).flatten(),axis=1)
    Q3[np.arange(T),ini]  =  1


    gB  =  2 * B @ (J * c @ E) @ (beta * Q2 * c @ E + B @ (np.linalg.matrix_power(I - J, 2) * c @ E)).T / beta_square
    B   += 0.1 * gB

    dif  =  np.max(np.absolute(Q3 - Q2))
    kcQ  =  k[ini]

    Q2   =  Q3
    loop += 1

该代码基于梯度下降算法,核心特征:

  • 矩阵B初始化为B = inv(I - J),通过B += 0.1 * gB迭代更新
  • J随稀疏矩阵Q2变化,Q2每列仅含一个1,生成逻辑为:
    ini =  np.argmax(c + (B @ (J * c) @ E).flatten(),axis=1)
    Q3[np.arange(T),ini]  =  1
    ...
    Q2  =  Q3
    

针对1000×1000规模矩阵运算,可从以下方向优化提速:

一、利用稀疏矩阵特性减少计算量

Q2是仅含单1的置换稀疏矩阵,J = beta * Q2同样稀疏,无需用稠密矩阵存储计算:

  • 使用scipy.sparse的csr_matrix或置换矩阵类存储Q2和J,避免稠密矩阵的冗余存储与运算
  • 用scipy.sparse.linalg.spsolve代替inv求解(I-J)B=I,避免直接求逆的高开销,稀疏线性代数工具对这类特殊矩阵的运算效率远高于稠密矩阵

二、简化矩阵运算维度,避免冗余操作

观察核心运算逻辑,可做如下简化:

  • J * c:因J是置换矩阵,可直接通过索引操作实现行置换,无需矩阵乘法
  • B @ (J * c) @ E:E是全1列向量,(J * c) @ E等价于对J*c行求和,可简化为np.sum(J * c, axis=1, keepdims=True),减少一次矩阵乘法
  • (B @ ...).flatten():直接保留二维数组形式参与argmax,无需额外展平操作

三、缓存重复计算项

循环内多次出现J * c @ E、beta * Q2 * c @ E等重复计算,提前缓存复用:

# 循环内提前缓存重复项
jc_E = (J * c) @ E
beta_q2c_E = beta * (Q2 * c) @ E
# 后续直接复用
ini = np.argmax(c + (B @ jc_E), axis=1)

四、优化稀疏矩阵生成

Q3无需初始化全零矩阵再赋值,直接用稀疏矩阵API创建:

from scipy.sparse import csr_matrix
Q3 = csr_matrix((np.ones(T), (np.arange(T), ini)), shape=(T, T))

五、借助编译与硬件加速

  • 用numba对循环内核心运算做JIT编译,将Python代码转为机器码,大幅提升计算速度
  • 使用MKL优化的NumPy版本(如Anaconda默认版本),MKL对矩阵乘法、求逆等操作有硬件加速支持

内容的提问来源于stack exchange,提问作者Zuba Tupaki

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 12:43:10