如何以广播方式批量判断二维数组A的行是否存在于二维数组B中
如何以广播方式批量判断二维数组A的行是否存在于二维数组B中
嘿,这个需求我之前处理类似问题时也遇到过,完全不用写循环,靠numpy的广播机制或者一些小技巧就能搞定!
先给你最直接的广播实现方案,完全贴合你想要的“无循环+广播”思路:
你之前用np.any(np.all(A[index] == B, axis=1))判断单行的逻辑是对的,要批量处理的话,只需要给A增加一个维度,让它能和B做广播比较就行:
import numpy as np # 先定义你的数组 A = np.array([[1, 1], [2, 2], [3, 1]]) B = np.array([[2, 1], [1, 2], [3, 1], [4, 2]]) # 核心代码:用广播实现批量判断 result = np.any(np.all(A[:, np.newaxis, :] == B, axis=2), axis=1) print(result) # 输出正好是你要的:array([False, False, True])
我来拆解下这段代码的逻辑:
A[:, np.newaxis, :]把A的形状从(3,2)改成了(3,1,2),相当于给每一行都套了一个“外层维度”- 这样和形状为
(4,2)的B比较时,numpy会自动把两者广播成(3,4,2)的数组——简单说就是让A的每一行都和B的所有行做逐元素比较 np.all(..., axis=2)是沿着最后一个维度(也就是每行的元素维度)做“全相等”判断,得到一个(3,4)的布尔数组,每个位置表示A的第i行和B的第j行是否完全匹配- 最后
np.any(..., axis=1)沿着B的行维度做“存在任意匹配”判断,就得到了A的每一行是否在B中的结果
如果你的数组行数特别多,广播生成的中间数组可能会占用较多内存,这时候可以用另一种更高效的技巧:把每行转换成一个单一的“虚拟标量”,再用np.isin判断,避免生成大的中间数组:
def check_rows_in_b(A, B): # 把每行的字节数据转换成一个统一的void类型,相当于把整行当成一个单独的元素 row_dtype = np.dtype((np.void, A.dtype.itemsize * A.shape[1])) # 转换成一维的view数组 A_rows = A.view(row_dtype).ravel() B_rows = B.view(row_dtype).ravel() # 直接用isin判断存在性 return np.isin(A_rows, B_rows) result = check_rows_in_b(A, B) print(result) # 同样输出:array([False, False, True])
这种方法的原理是把每一行的二进制数据直接打包成一个不可分割的类型,这样二维数组就变成了一维数组,判断起来不仅更快,内存占用也小很多,适合处理大规模数据。
两种方法都能满足你的无循环需求,选哪种看你的数组规模和代码可读性要求就好~
备注:内容来源于stack exchange,提问作者RCTW
相关产品推荐
相关产品推荐

