计算区分度:加速矩阵行所有组合的运算
加速二进制矩阵测试用例区分度计算的优化方案
嘿,这个问题我太熟了——你的原代码思路完全正确,但在个体数量较多时,itertools.combinations带来的O(N²)遍历会让速度慢到离谱。咱们直接用数学公式把它改成纯向量化操作,性能能提升好几个数量级!
为什么原代码慢?
原代码遍历所有个体对(共C(N,2) = N*(N-1)/2对),逐对比较每一列的差异再累加。当N=1000时,就有499500次循环;N=10000时,这个数字会飙升到近5000万——完全没必要做这么多重复计算。
核心数学原理(二进制矩阵专属)
对于二进制矩阵的某一列来说:
- 假设该列有
k个1,N-k个0(N是个体总数) - 能区分的个体对数量 = k(N-k)*
原因很简单:只有一个取1、一个取0的个体对才会被该列区分,而两个1或两个0的对无法区分。所有这样的组合数正好是1的数量乘以0的数量,和你遍历所有对统计的结果完全一致。
优化后的代码
用NumPy的向量化操作直接计算,全程没有循环:
import numpy as np # 假设G是你的二进制NumPy矩阵(元素为0或1) num_individuals = G.shape[0] # 计算每列中1的数量 count_ones_per_col = G.sum(axis=0) # 直接用公式计算区分度 distinctions = count_ones_per_col * (num_individuals - count_ones_per_col)
验证正确性
举个小例子测试:
G = np.array([ [1, 0, 1], [0, 0, 1], [1, 1, 0] ])
- 原代码计算结果:
array([2, 2, 2]) - 优化后代码计算结果:
count_ones_per_col = [2,1,2],distinctions = 2*(3-2) = 2,1*(3-1)=2,2*(3-2)=2,完全一致。
额外说明
如果你的矩阵不是严格的0/1二进制(比如存在其他值),可以先转成标准二进制:
# 将非0值转为1,0保持不变 G_binary = (G != 0).astype(int)
这个优化后的方案时间复杂度是O(NM),相比原代码的O(N²M),当N较大时速度提升极其明显——比如N=1000时,速度是原代码的500倍左右;N=10000时,能快5000倍以上,完全能满足性能需求。
内容的提问来源于stack exchange,提问作者voo
相关产品推荐
相关产品推荐

