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

scipy中csr_matrix切片A[index,:]为何极快?如何实现高性能CSR矩阵切片

scipy CSR矩阵行切片高性能的核心原理

CSR格式稀疏矩阵的存储由三个连续数组构成:

  • indptr:长度为行数+1的行偏移数组,第i行的所有非零元对应indices[indptr[i]:indptr[i+1]]和data[indptr[i]:indptr[i+1]]区间
  • indices:存储每个非零元的列索引
  • data:存储每个非零元的数值

scipy原生的A[index,:]行切片完全没有走稀疏矩阵乘法逻辑,而是针对行切片的场景做了定向优化,全程只有顺序内存操作,步骤为:

  1. 遍历传入的行索引列表index,累加每个选中行的非零元数量,直接生成新矩阵的indptr数组,这一步复杂度为O(k),k为选中的行数
  2. 根据indptr最后一位的值(即新矩阵总非零元数)一次性分配好新矩阵的indices和data数组内存
  3. 遍历每个选中行,把原矩阵对应行的indices、data段直接整块拷贝到新数组的对应位置,这一步复杂度为O(nnz_selected),即选中行包含的总非零元数

整个过程没有乘加计算、没有列索引排序/去重、没有动态内存扩容,缓存命中率极高,所以速度非常快。

你之前测试的构造提取矩阵做乘法的方案慢6~7倍是完全正常的:这个方案走的是通用稀疏矩阵乘法(SpGEMM)逻辑,哪怕提取矩阵每行只有一个值为1的非零元,依然要执行完整的乘法流程,包括中间结果累加、列索引去重排序、多次内存申请释放,引入了大量完全不必要的开销。

Eigen 3.3.8实现同性能行切片的方案

不要使用Eigen自带的逐行插入、稀疏乘法、通用块截取接口,这些接口都没有针对行切片场景做优化,直接操作CSR矩阵的底层存储数组,复现scipy的实现逻辑即可达到同等甚至更高的性能,具体实现步骤如下:

  1. 直接获取原矩阵的三个底层数组指针:outerIndexPtr()对应indptr,innerIndexPtr()对应列索引数组,valuePtr()对应数值数组
  2. 遍历选中行索引列表,预计算新矩阵的indptr数组,一次性算出总非零元数量
  3. 给新矩阵预分配刚好足够的内存,避免动态扩容开销
  4. 逐行将原矩阵对应行的列索引、数值整块拷贝到新矩阵的对应位置,不要逐个非零元调用插入接口
  5. 给新矩阵设置正确的非零元计数即可

参考实现代码:

#include <Eigen/Sparse>
#include <cstring>
#include <vector>

// 定义行优先存储的CSR矩阵类型,和scipy csr_matrix存储格式对齐
using SpMatCSR = Eigen::SparseMatrix<double, Eigen::RowMajor>;

SpMatCSR rowSlice(const SpMatCSR& A, const std::vector<int>& selected_rows) {
    const int new_row_cnt = selected_rows.size();
    const int col_cnt = A.cols();
    const int* old_indptr = A.outerIndexPtr();
    const int* old_col_idx = A.innerIndexPtr();
    const double* old_val = A.valuePtr();

    // 计算新矩阵的行偏移数组
    std::vector<int> new_indptr(new_row_cnt + 1, 0);
    for (int i = 0; i < new_row_cnt; ++i) {
        int r = selected_rows[i];
        new_indptr[i+1] = new_indptr[i] + (old_indptr[r+1] - old_indptr[r]);
    }
    const int total_nnz = new_indptr.back();

    SpMatCSR res(new_row_cnt, col_cnt);
    res.makeCompressed();
    res.resizeNonZeros(total_nnz);
    int* new_indptr_ptr = res.outerIndexPtr();
    int* new_col_idx = res.innerIndexPtr();
    double* new_val = res.valuePtr();

    // 拷贝行偏移
    memcpy(new_indptr_ptr, new_indptr.data(), sizeof(int) * (new_row_cnt + 1));
    // 逐行拷贝非零元的列索引和数值
    for (int i = 0; i < new_row_cnt; ++i) {
        int r = selected_rows[i];
        const int old_start = old_indptr[r];
        const int nnz_r = old_indptr[r+1] - old_start;
        const int new_start = new_indptr[i];
        memcpy(new_col_idx + new_start, old_col_idx + old_start, sizeof(int) * nnz_r);
        memcpy(new_val + new_start, old_val + old_start, sizeof(double) * nnz_r);
    }

    return res;
}

实现注意事项

  • 必须使用行优先(Eigen::RowMajor)的稀疏矩阵类型,列优先的是CSC格式,行切片效率极低
  • 不要用res.insert()逐元素插入非零元,该接口每次插入都会做内存边界检查和位置调整,开销是整块拷贝的数十倍
  • 提前一次性分配好所有需要的内存,避免切片过程中动态扩容
  • 该实现和scipy原生切片行为完全一致,支持重复行索引、乱序行索引,时间复杂度和scipy实现完全对齐,性能基本没有差距。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 09:15:40