高难度索引挑战:如何用单行索引实现成对索引计数?
兄弟,这个需求我太熟了——当数据量上去之后,逐行遍历统计每对索引绝对会慢到让人抓狂,用矩阵化的索引操作才是提升效率的核心!下面我给你拆解下最实用的思路和实现步骤,保证效率拉满。
第一步:将原始数据转换为「存在矩阵」
首先我们要把每行的K个索引,转换成一个二进制存在矩阵:X[row_idx, idx] = 1表示该行包含索引idx,否则为0。这一步是所有高效统计的基础,完全不用循环,用向量索引赋值就能搞定。
举个Python numpy的实现例子:
假设你的原始数据是形状为(N, K)的数组data,INDEX_SIZE是索引的总数量(比如索引范围0~M,那INDEX_SIZE = M+1):
import numpy as np # 初始化存在矩阵,用int8省内存,N行×INDEX_SIZE列,全0 X = np.zeros((N, INDEX_SIZE), dtype=np.int8) # 生成行号和列号的索引数组,批量赋值1 rows = np.repeat(np.arange(N), K) # 每行重复K次,对应该行的K个元素 cols = data.flatten() # 把所有索引展开成一维数组 X[rows, cols] = 1
如果某行里有重复索引(比如同一个索引出现多次),记得先对每行去重,因为我们只关心「是否出现」,而非出现次数:
# 对每行去重,得到每行的唯一索引集合 data_unique = np.array([np.unique(row) for row in data]) # 重新生成行号和列号 rows = np.concatenate([np.full(len(row), i) for i, row in enumerate(data_unique)]) cols = np.concatenate(data_unique) # 再初始化并填充存在矩阵 X = np.zeros((N, INDEX_SIZE), dtype=np.int8) X[rows, cols] = 1
第二步:用矩阵运算直接计算四类统计量
有了存在矩阵X,我们可以用高度优化的矩阵运算直接算出四个统计量,完全不用遍历每对索引:
1. n11[i,j]:i和j共同出现的行数
这其实就是存在矩阵的转置乘原矩阵——X.T @ X的[i,j]位置值,就是所有行中i和j同时出现的次数(本质是两个列向量的点积):
n11 = X.T @ X
numpy的矩阵乘法用的是底层BLAS/LAPACK优化,速度快到离谱。
2. n10[i,j]:i出现但j未出现的行数
先算出每个索引出现的总行数count_i,然后用i的总出现行数减去i和j共同出现的行数即可:
count_i = X.sum(axis=0, keepdims=True) # 形状(1, INDEX_SIZE),每个索引的总出现行数 n10 = count_i.T - n11 # 广播运算,自动对应每个i,j对
3. n01[i,j]:i未出现但j出现的行数
和上面逻辑类似,用j的总出现行数减去i和j共同出现的行数:
count_j = X.sum(axis=0, keepdims=True) n01 = count_j - n11
4. n00[i,j]:i和j均未出现的行数
用容斥原理计算:总行数N减去i出现的行数,减去j出现的行数,再加回i和j共同出现的行数(因为前面减重复了):
n00 = N - count_i.T - count_j + n11
超大索引范围的优化方案
如果INDEX_SIZE特别大(比如超过1e5),普通的密集矩阵会占太多内存,这时候可以用稀疏矩阵来存储,scipy的稀疏矩阵支持高效的矩阵运算:
from scipy.sparse import csr_matrix # 构建稀疏存在矩阵 X_sparse = csr_matrix((np.ones(len(cols), dtype=np.int8), (rows, cols)), shape=(N, INDEX_SIZE)) # 稀疏矩阵乘法计算n11 n11_sparse = X_sparse.T @ X_sparse # 后续的n10/n01/n00可以把稀疏矩阵转成密集数组计算,或者继续用稀疏运算 count_i_sparse = X_sparse.sum(axis=0).A # 转成密集数组
为什么这种方式效率高?
- 彻底避免了O(NINDEX_SIZE²)的暴力遍历,时间复杂度降到O(NK + INDEX_SIZE²),其中矩阵乘法是CPU并行优化的,速度提升几个数量级。
- 用int8或稀疏矩阵存储存在矩阵,大幅降低内存占用,避免内存溢出。
内容的提问来源于stack exchange,提问作者user1581390

