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

PyTorch/Numpy中小向量与超大稀疏矩阵乘法的优化咨询

针对小向量a与超大稀疏矩阵b的乘法优化需求(b只读、需多次和不同a计算),以下是基于稀疏特性的实用优化方案,附代码示例:

方案1:转换为CSC稀疏矩阵(开发效率最高)

因为a @ b本质是对b的每一列做加权求和,而CSC(压缩稀疏列)格式天然适配列操作。scipy.sparse的CSC矩阵会自动跳过零元素,只计算非零项的贡献,预处理仅需一次,后续计算速度远快于密集矩阵乘法。

import numpy as np
from scipy.sparse import csc_matrix

# 从memmap加载原始b矩阵(假设已提前保存为memmap文件)
B = 32
M = 10000000
b_memmap = np.memmap('b_matrix.npy', dtype=bool, mode='r', shape=(B, M))

# 预处理:转换为CSC格式(仅执行一次)
b_csc = csc_matrix(b_memmap)

# 多次执行与不同a的乘法
for _ in range(100):
    a = np.random.rand(B)
    result = a @ b_csc  # 自动利用稀疏性优化计算
    # 处理结果...

方案2:自定义稀疏存储(针对极小B的轻量优化)

当B很小(比如你的场景是32),可以手动提取每一列的非零索引,用更紧凑的格式存储,再通过Numba JIT加速循环计算,比scipy.sparse的开销更低。

import numpy as np
from numba import njit

B = 32
M = 10000000

# 从memmap加载b
b_memmap = np.memmap('b_matrix.npy', dtype=bool, mode='r', shape=(B, M))

# 预处理:提取非零索引的紧凑存储(仅执行一次)
# 用indptr记录每一列非零元素的起始位置,indices存储所有非零元素的行索引
counts = np.array([np.sum(b_memmap[:, j]) for j in range(M)], dtype=np.int32)
indptr = np.zeros(M + 1, dtype=np.int64)
indptr[1:] = np.cumsum(counts)
indices = np.zeros(indptr[-1], dtype=np.int32)

pos = 0
for j in range(M):
    cols = np.where(b_memmap[:, j])[0]
    indices[pos:pos + len(cols)] = cols
    pos += len(cols)

# Numba JIT并行加速计算
@njit(parallel=True)
def compute_result(a, indices, indptr, M):
    result = np.zeros(M, dtype=np.float64)
    for j in range(M):
        start = indptr[j]
        end = indptr[j + 1]
        total = 0.0
        for k in range(start, end):
            total += a[indices[k]]
        result[j] = total
    return result

# 多次计算
for _ in range(100):
    a = np.random.rand(B)
    result = compute_result(a, indices, indptr, M)
    # 处理结果...

方案3:位掩码+向量指令(极致性能优化)

因为B=32刚好匹配32位整数的位数,可以把每一列的bool值打包成一个32位整数(位掩码),内存占用极低(仅40MB),再通过Numba编译出利用CPU向量指令的代码,速度达到最优。

import numpy as np
from numba import njit, uint32, float64

B = 32
M = 10000000

# 从memmap加载b
b_memmap = np.memmap('b_matrix.npy', dtype=bool, mode='r', shape=(B, M))

# 预处理:将每一列转换为32位位掩码(仅执行一次)
bitmasks = np.zeros(M, dtype=np.uint32)
for j in range(M):
    mask = 0
    for i in range(B):
        if b_memmap[i, j]:
            mask |= (1 << i)
    bitmasks[j] = mask

# Numba JIT并行加速,利用位运算快速求和
@njit(parallel=True)
def compute_with_bitmask(a, bitmasks, B, M):
    result = np.zeros(M, dtype=np.float64)
    for j in range(M):
        mask = bitmasks[j]
        total = 0.0
        i = 0
        while mask:
            if mask & 1:
                total += a[i]
            mask >>= 1
            i += 1
        result[j] = total
    return result

# 多次计算
for _ in range(100):
    a = np.random.rand(B)
    result = compute_with_bitmask(a, bitmasks, B, M)
    # 处理结果...

方案选择建议

  • 追求开发效率:直接用方案1,代码最简,无需手动处理稀疏逻辑。
  • 追求极致性能:选方案3,内存占用最少,速度最快,完美适配B=32的场景;方案2适合B稍大的情况。
  • 所有预处理步骤仅需执行一次,后续和不同a的计算都能复用预处理后的结构,大幅降低重复计算量。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 11:25:14