numpy.intersect1d无法检测全部交集元素的原因排查
问题描述
我有如下两个数组,希望找出所有交集元素及对应索引,遂使用numpy.intersect1d()方法,但得到的结果中缺失了3.5e-04和1.5e-01这两个元素。使用集合求交集的方式也出现同样问题,代码示例如下:
import numpy as np array1 = np.array([1.1, 1.5, 3.5, 6.5, 10, 15, 35, 65, 100, 150, 350, 650, 1000, 1500, 3500]) * 10**-4 array2 = np.array([6.5, 10, 15, 35, 65, 100, 150, 350, 650, 1000, 1500, 3500, 6500, 10000, 15000, 35000, 65000]) * 10**-5 intersected, i_1, i_2 = np.intersect1d(array1, array2, return_indices=True)
intersected结果为:
[1.5e-04 6.5e-04 1.0e-03 1.5e-03 3.5e-03 6.5e-03 1.0e-02 1.5e-02 3.5e-02 6.5e-02 1.0e-01 3.5e-01]
使用集合求交集:
np.array(list(set(array1) & set(array2)))
结果同样缺失上述元素,请问这是为什么?
原因分析
核心问题出在浮点数的精度误差上:
- 计算机以二进制格式存储浮点数,绝大多数十进制小数无法被精确表示,只能存储近似值。这会导致两个在十进制下看起来完全相等的数值,在内存中的实际存储值存在微小的差异。
- 以你提到的
3.5e-04为例:array1中的该值是3.5 * 10**-4,实际存储的可能是0.00035000000000000003这类近似值;array2中对应的35 * 10**-5,实际存储的可能是0.00034999999999999996这类近似值;
numpy.intersect1d和集合的交集判断都是基于严格相等(==),这种微小的差异会被判定为不相等,因此这两个元素没有被纳入交集结果。
你可以通过打印高精度数值来验证这一点:
print("{0:.20f}".format(array1[2])) # 输出array1中3.5e-04的高精度表示 print("{0:.20f}".format(array2[3])) # 输出array2中对应值的高精度表示
解决方法
不要依赖严格相等判断,改用近似相等的方式来处理:
- 使用
np.isclose设定精度阈值判断
可以通过广播机制检查array1中的每个元素是否在array2中有近似相等的值,示例代码:# 设定精度阈值,比如1e-10 threshold = 1e-10 # 生成布尔掩码,标记array1中哪些元素在array2中有近似匹配 mask = np.any(np.isclose(array1[:, np.newaxis], array2, atol=threshold), axis=1) # 获取交集元素 intersected_values = array1[mask] # 获取对应索引 indices_array1 = np.where(mask)[0] indices_array2 = [np.where(np.isclose(val, array2, atol=threshold))[0][0] for val in intersected_values] - 对数值进行四舍五入
如果你的数值有固定的有效位数,可以先对数组四舍五入到指定小数位,再求交集:
注意:这种方法可能会误判接近阈值的数值,需要根据业务场景选择合适的小数位数。# 四舍五入到8位小数 rounded_array1 = np.round(array1, 8) rounded_array2 = np.round(array2, 8) intersected, i_1, i_2 = np.intersect1d(rounded_array1, rounded_array2, return_indices=True)
内容的提问来源于stack exchange,提问作者Olymp 007
相关产品推荐
相关产品推荐

