如何获取SciPy稀疏矩阵中阈值以下值的索引?解决索引匹配问题
如何过滤SciPy稀疏数组并获取符合条件值的索引
问题核心在于DOK转CSC时的重复索引合并规则,以及验证方式的误区。以下是正确的解决方法:
正确获取索引的步骤
CSC格式稀疏数组的data与nonzero()返回的行/列索引是严格一一对应的,直接用掩码过滤即可:
import numpy as np from scipy.sparse import dok_array, csc_array # 生成无重复索引的DOK数组,避免重复赋值导致值被求和 pairs = list(set(zip(np.random.randint(0, 1000, 100), np.random.randint(0, 1000, 100)))) m = dok_array((1000, 1000)) for i, j in pairs: m[i, j] = np.random.random() threshold = 0.3 tmp = csc_array(m) # 生成掩码并过滤得到符合条件的索引 mask = tmp.data < threshold rows, cols = tmp.nonzero() filtered_rows = rows[mask] filtered_cols = cols[mask] # 可靠的验证方式:直接检查被掩码选中的data值 assert np.all(tmp.data[mask] < threshold), "验证通过"
原代码断言失败的原因
- 重复索引求和:如果生成的
(i,j)对存在重复,DOK转CSC时会自动将同一位置的所有值求和,导致tmp.data中的值是求和结果而非最后一次赋值的值。比如同一位置两次赋值0.2,求和后变为0.4,会被掩码排除,与预期不符。 - 验证逻辑问题:
tmp[i_mask, j_mask]返回的稠密数组可能受重复索引或浮点数精度影响,直接检查tmp.data[mask]才是最准确的验证方式。
从CSC内部属性直接生成索引
若不想使用nonzero(),可通过CSC的indptr属性生成列索引:
# 从indptr生成列索引数组 cols = np.repeat(np.arange(tmp.shape[1]), np.diff(tmp.indptr)) filtered_rows = tmp.indices[mask] filtered_cols = cols[mask]
内容的提问来源于stack exchange,提问作者Steven C. Howell
相关产品推荐
相关产品推荐

