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

如何优化Cython中动态稀疏矩阵的列表实现性能?

动态稀疏二进制矩阵的Cython优化问题

我需要实现一个难以向量化的函数,因此采用Cython处理嵌套循环。该问题复杂度类似朴素矩阵乘法,但至少有一个矩阵为二进制且通常高度稀疏,且该稀疏矩阵会在程序运行过程中动态变化。

最初选择嵌套列表格式(存储所有非零索引)实现稀疏版spgreedy函数,但发现它比密集版greedy函数慢近一个数量级。我推测性能瓶颈源于Python列表的使用,但由于需要对内部列表执行remove和append操作,认为CSR这类静态稀疏格式效率更低,查阅相关问题后也未找到适配动态增删场景的方案。

稀疏版与密集版函数代码

import numpy as np
cimport numpy as np
from libc.math cimport exp

def spgreedy(np.ndarray[double, ndim=2] J, 
             np.ndarray[double, ndim=1] h,
             list[list[int]] S,
             double temp,
             ):

    cdef int n, d 
    n = len(S)
    d = J.shape[0]

    cdef int j, k
    cdef double dot, curr, prob

    for s in S:
        for j in range(d):

            if j in s:
                s.remove(j)

            dot = h[j]
            for k in s:
                dot += J[j, k]

            curr = dot / temp

            if curr < -100:
                prob = 0.0
            elif curr > 100:
                prob = 1.0
            else:
                prob = 1.0 / (1.0 + exp(-curr))

            if np.random.rand() < prob:
                s.append(j)

    return S

def greedy(np.ndarray[double, ndim=2] J, 
           np.ndarray[double, ndim=1] h,
           np.ndarray[int, ndim=2] S,
           double temp,
           ):

    cdef int n, d 
    n = len(S)
    d = J.shape[0]

    cdef int i, j, k
    cdef double dot, curr, prob

    for i in range(n):
        for j in range(d):

            dot = h[j]
            for k in range(d):
                dot += J[j, k] * S[i,k]

            curr = dot / temp

            if curr < -100:
                prob = 0.0
            elif curr > 100:
                prob = 1.0
            else:
                prob = 1.0 / (1.0 + exp(-curr))

            S[i,j] = 1*(np.random.rand() < prob)

    return S

测试代码及结果

import time
import numpy as np

d = 1000
n = 50

J = 1.0*np.random.choice([-1,0,1], size=(d,d), p=[0.05,0.9,0.05])
J = np.triu(J,1) + np.triu(J,1).T
h = -np.ones(d)

S = np.random.choice([0,1], size=(n, d), p=[0.95, 0.05])
Slist = [np.where(s)[0].tolist() for s in S]

t0 = time.time()
_ = greedy(J, h, S, 1e-5)
print(f"dense: {time.time() - t0}")

t0 = time.time()
_ = spgreedy(J, h, Slist, 1e-5)
print(f"sparse: {time.time() - t0}")

运行结果:

dense: 0.04369068145751953
sparse: 0.18976712226867676

我既不想在密集版内层for k in range(d)循环中浪费算力(大部分元素无贡献),也想解决列表实现的性能问题,希望找到更合适的数据结构或优化方式。


内容的提问来源于stack exchange,提问作者Matteo Alleman

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 01:43:25