如何在不转换为稠密矩阵的前提下为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
相关产品推荐
相关产品推荐

