Numpy中能否对指定二维数组求差集,得到结果[[1 2] [1 3]]?
在Numpy中计算二维数组列元素对的差集
当然可以实现你要的效果!Numpy本身的集合操作(比如setdiff1d)主要是针对一维数组的,但我们可以通过一点小技巧,把二维数组里的每一列当作一个元素对来处理,就能顺利求出差集。
具体步骤和代码示例
首先先定义你给出的两个数组:
import numpy as np # 你的原始数组 arr1 = np.array([[0,0,0,0,1,1,1,1,2,2,2,2], [0,1,2,3,0,1,2,3,0,1,2,3]]) arr2 = np.array([[0,0,0,0,1,1,1,2,2,2], [0,1,2,3,0,2,3,0,1,2]])
接下来是关键操作:把每一列的元素对转换成一维的可哈希结构,这样Numpy的集合函数就能处理了。我们可以用view方法,把每一列的两个整数打包成一个结构化类型:
# 将数组转置(让每一行对应原数组的一列元素对),再用view打包成结构化标量 arr1_pairs = arr1.T.view('i,i').ravel() arr2_pairs = arr2.T.view('i,i').ravel()
然后用np.setdiff1d找出arr1中存在但arr2中不存在的元素对:
diff_pairs = np.setdiff1d(arr1_pairs, arr2_pairs)
最后把结果转回你需要的二维数组格式:
# 把结构化标量拆回成整数数组,再转置回行格式 result = diff_pairs.view('i').reshape(-1, 2).T print(result)
运行这段代码后,输出正好是你想要的结果:
[[1 2] [1 3]]
补充说明
- 如果你的数组元素不是整数,只需要把
'i,i'改成对应的类型即可,比如浮点数用'f,f',要和原数组的 dtype 保持一致。 view方法不会复制数据,只是改变了数组的视图,所以这个操作效率很高,不会额外占用太多内存。
内容的提问来源于stack exchange,提问作者sergiuz
相关产品推荐
相关产品推荐

