求解二进制Numpy矩阵按最高公共列统计1配对数的高效算法
优于O(n²k)的实现方案
因为你的场景中n远大于k,我们可以利用二进制位运算特性,将时间复杂度降到O(nk + 4^k),k较小时性能远高于原有实现。
核心逻辑
两个行的最高公共列,等价于把每行按列索引权重(列j对应权重2^j)编码为整数后,两个整数按位与结果的最高置位索引。我们只需统计不同编码的出现频次,再遍历所有编码对计算贡献即可:
- 每行编码:将每一行的二进制数组按权重累加为整数,统计每个整数的出现频次
- 遍历所有编码对,计算按位与的最高置位索引,将对应配对数累加到结果数组的对应位置
实现代码
import numpy as np def get_num_pairs_fast(bin_mat): n, k = bin_mat.shape # 列j对应权重2^j,最高列对应最高位 powers = 1 << np.arange(k) row_codes = bin_mat.dot(powers) # 统计唯一编码的出现次数 unique_codes, counts = np.unique(row_codes, return_counts=True) m = len(unique_codes) ans = np.zeros(k, dtype=np.int32) for i in range(m): x = unique_codes[i] cnt_x = counts[i] # 相同编码的内部配对 if cnt_x >= 2: pairs = cnt_x * (cnt_x - 1) // 2 highest_bit = x.bit_length() - 1 ans[highest_bit] += pairs # 不同编码的配对 for j in range(i + 1, m): y = unique_codes[j] cnt_y = counts[j] and_val = x & y pairs = cnt_x * cnt_y highest_bit = and_val.bit_length() - 1 ans[highest_bit] += pairs return ans
验证结果
测试用例1
输入:
arr = np.array([[1, 1, 1, 0, 1, 0], [1, 0, 1, 1, 1, 0], [1, 0, 1, 0, 0, 0], [1, 1, 1, 1, 1, 0], [1, 0, 1, 0, 1, 1], [1, 0, 1, 0, 1, 1], [1, 1, 1, 1, 0, 0], [1, 1, 1, 0, 1, 0]], dtype=np.int32)
输出:array([ 0, 0, 11, 2, 14, 1], dtype=int32),和预期完全一致。
测试用例2
输入:
arr = np.array([[1, 1, 0], [1, 1, 0], [1, 1, 0], [1, 1, 0], [1, 0, 1], [1, 0, 1], [1, 0, 1], [1, 0, 1]], dtype=np.int32)
输出:array([16, 6, 6], dtype=int32),和预期完全一致。
复杂度说明
- 编码阶段:O(nk),仅需要遍历矩阵所有元素一次
- 统计阶段:O(m²),m是唯一编码的数量,最多为2k,最差复杂度为O(4k)
当k<=12时,412=16777216,可在毫秒级完成,n哪怕到百万级也仅需要千万级操作,性能远高于原O(n²k)实现。如果k更大(k>15),可以用基于快速Zeta变换的容斥方法,将统计阶段复杂度降到O(k*2k)进一步提升性能。
内容的提问来源于stack exchange,提问作者BadBayesian
相关产品推荐
相关产品推荐

