如何在Scipy稀疏矩阵中选取指定索引对应的行?切片方法报错
在Scipy稀疏矩阵中选取指定索引行的正确姿势
嘿,你碰到的这个报错其实是因为scipy.sparse.eye()默认生成的是dia_matrix(对角稀疏矩阵),这种类型确实不支持像普通NumPy数组那样直接用[:2, :]做下标切片。别担心,有好几种办法能解决这个问题:
方法1:转成支持切片的稀疏矩阵类型
最省心的方式是把dia_matrix转换成csr_matrix或者csc_matrix——这俩是Scipy稀疏矩阵里最常用、也支持切片操作的类型。转完之后就能像用NumPy数组一样切片了:
import scipy.sparse as sparse # 生成对角矩阵并转为csr格式 sparse_mat = sparse.eye(3).tocsr() # 轻松选取前2行 selected_rows = sparse_mat[:2, :] print(selected_rows.toarray())
方法2:用矩阵自带的getrow()/getrows()方法
如果不想转换矩阵类型,也可以用稀疏矩阵原生的方法来挑行:
getrow(i):单独选取第i行getrows(indices):一次性选取多个指定索引的行,参数传索引列表就行
举个例子:
import scipy.sparse as sparse dia_mat = sparse.eye(3) # 选取第0和第1行 selected_rows = dia_mat.getrows([0, 1]) print(selected_rows.toarray())
方法3:用索引矩阵实现复杂筛选
要是你需要更灵活的行筛选(比如非连续的、带条件的索引),还可以构建一个行选择矩阵,和原矩阵相乘来实现:
import scipy.sparse as sparse dia_mat = sparse.eye(3) # 定义要选取的行索引 target_indices = [0, 1] # 构建行选择的稀疏矩阵 row_selector = sparse.csr_matrix( ([1]*len(target_indices), (range(len(target_indices)), target_indices)), shape=(len(target_indices), dia_mat.shape[0]) ) # 矩阵乘法得到选中的行 selected_rows = row_selector @ dia_mat print(selected_rows.toarray())
总的来说,第一种方法(转csr/csc后切片)最符合我们平时用NumPy的习惯,上手最快。
内容的提问来源于stack exchange,提问作者David Masip
相关产品推荐
相关产品推荐

