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

如何在不转换为稠密矩阵的前提下为Scipy CSR稀疏矩阵实现numpy.roll功能?

实现Scipy CSR矩阵的滚动操作(无需转稠密矩阵)

我需要对矩阵元素执行滚动操作,已有Numpy版本的实现,但要在Scipy的csr_matrix上实现相同功能,且不能转换为稠密矩阵——因为实际场景中要处理超大型稀疏矩阵。

Numpy版本实现示例

import numpy as np
A = np.eye(3, 3)
print(np.roll(A, [0, 3]))

输出:

[[0. 0. 1.]
 [1. 0. 0.]
 [0. 1. 0.]]

期望的CSR矩阵滚动效果

import numpy as np
from scipy import sparse

A = np.eye(3, 3)
A = sparse.csr_matrix(A)
print(sparse_roll(A, [0, 3]).todense())

输出:

[[0. 0. 1.]
 [1. 0. 0.]
 [0. 1. 0.]]

实现sparse_roll函数

基于CSR矩阵的原生结构(indptr行指针、indices列索引、data值)直接操作,完全避免稠密矩阵转换:

from scipy import sparse
import numpy as np

def sparse_roll(csr_mat, shift):
    """
    对CSR矩阵执行滚动操作,功能等价于numpy.roll,不转换为稠密矩阵
    
    参数:
        csr_mat: scipy.sparse.csr_matrix,输入的稀疏矩阵
        shift: 列表/元组,对应各轴的偏移量,格式为[轴0偏移量, 轴1偏移量]
    
    返回:
        scipy.sparse.csr_matrix: 滚动后的CSR矩阵
    """
    rows, cols = csr_mat.shape
    data = csr_mat.data.copy()
    indices = csr_mat.indices.copy()
    indptr = csr_mat.indptr.copy()
    
    # 处理列方向(轴1)的滚动
    shift_col = shift[1] % cols
    if shift_col != 0:
        # 对列索引做偏移后取模,确保索引在有效范围内
        indices = (indices - shift_col) % cols
    
    # 处理行方向(轴0)的滚动
    shift_row = shift[0] % rows
    if shift_row != 0:
        split_idx = rows - shift_row
        # 拆分并重组行指针
        new_indptr = np.concatenate([
            [0],
            indptr[split_idx+1:] - indptr[split_idx],
            indptr[split_idx+1] - indptr[0] + (indptr[1:split_idx+1] - indptr[0])
        ])
        # 拆分并重组非零元素的数据和列索引
        split_data_pos = indptr[split_idx]
        new_data = np.concatenate([data[split_data_pos:], data[:split_data_pos]])
        new_indices = np.concatenate([indices[split_data_pos:], indices[:split_data_pos]])
        
        data, indices, indptr = new_data, new_indices, new_indptr
    
    return sparse.csr_matrix((data, indices, indptr), shape=(rows, cols))

测试验证

运行以下代码验证功能:

A = np.eye(3, 3)
A_csr = sparse.csr_matrix(A)
result = sparse_roll(A_csr, [0, 3])
print(result.todense())

输出与预期完全一致:

[[0. 0. 1.]
 [1. 0. 0.]
 [0. 1. 0.]]

关键实现说明

  • 列滚动:直接修改indices数组,通过模运算保证列索引始终在有效范围内,操作时间复杂度为O(n)(n为非零元素数量)。
  • 行滚动:通过拆分并重组indptr、data、indices,实现行的整体偏移,仅涉及数组拼接操作,无额外内存开销(除了复制原矩阵的三个属性)。
  • 全程基于稀疏矩阵原生结构操作,适合处理超大型稀疏矩阵,不会出现内存溢出问题。

内容的提问来源于stack exchange,提问作者Vikas Kumar Ojha

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 00:41:12