NumPy中浮点数与整数数组的匹配行为不一致问题
问题:NumPy子数组完全匹配判断错误
整数数组场景(行为符合预期)
以下代码用于判断array1的每个子数组是否在array2中存在,运行结果符合预期:
import numpy as np # 创建5x3数组 array1 = np.array([[1, 2, 3], [4, 5, 6], [7, 8, 9], [101, 110, 120], [13, 14, 15]]) # 创建8x3数组 array2 = np.array([[1, 2, 3], [7, 8, 9], [4, 5, 6], [16, 17, 18], [19, 20, 21], [10, 11, 12], [22, 23, 24], [13, 14, 15]]) # 检查array1每个元素与array2所有元素的相等性 result = np.any(array1[:, None] == array2, axis=2) # 转换为nx1布尔矩阵 final_result = np.squeeze(np.any(result, axis=1, keepdims=True))
输出:
[ True True True False True]
浮点数数组场景(行为异常)
但切换到浮点数数组时,代码逻辑出现问题:
import numpy as np array3 = np.array([[90., -40.8 , -1.35],[0., 0., -1.35],[0., 0.,-1.35], [100., 0.,-1.35], [-10., 50.8,-1.35], [100., -50.8,5.], [ 29.5, 0., 5. ],[ -2.89, -50.8,-1.35], [0., 0.,0.], [0., 0.,0.], [0., 0.,0.], [0., 0.,0.], [0., 0.,0.]], dtype=float) array4 = np.array([[90, -50.8, -1.35], [90, -50.8, 5], [90, -50.8, -1.35], [-10, 50.8, -1.35], [-10, 50.8, 5], [90, -50.8, 5]], dtype=float) # 检查array4每个元素与array3所有元素的相等性 result = np.any(array4[:, None] == array3, axis=2) # 转换为nx1布尔矩阵 final_result = np.squeeze(np.any(result, axis=1, keepdims=True))
输出:
[ True True True True True True]
实际仅array4[3]的子数组在array3中完全存在,其余均不匹配,预期输出应为:
[False False False True False False]
问题原因
原代码中使用了两次np.any():
- 第一次
np.any(..., axis=2):只要子数组中有任意一个元素匹配就返回True - 第二次
np.any(..., axis=1):只要array3中有任意一行满足上述条件就返回True
这导致只要子数组有一个元素和array3中某行的对应元素相等,就会判定为匹配,而非整个子数组完全匹配。
解决方案
要实现整个子数组所有元素完全匹配的判断,需要:
- 先通过
np.all(..., axis=2)检查每行的所有元素是否完全相等(即子数组完全匹配) - 再通过
np.any(..., axis=1)检查array3中是否存在这样的完全匹配行
修改后的代码
import numpy as np array3 = np.array([[90., -40.8 , -1.35],[0., 0., -1.35],[0., 0.,-1.35], [100., 0.,-1.35], [-10., 50.8,-1.35], [100., -50.8,5.], [ 29.5, 0., 5. ],[ -2.89, -50.8,-1.35], [0., 0.,0.], [0., 0.,0.], [0., 0.,0.], [0., 0.,0.], [0., 0.,0.]], dtype=float) array4 = np.array([[90, -50.8, -1.35], [90, -50.8, 5], [90, -50.8, -1.35], [-10, 50.8, -1.35], [-10, 50.8, 5], [90, -50.8, 5]], dtype=float) # 先检查每行所有元素是否完全匹配,再检查是否存在匹配行 result = np.all(array4[:, None] == array3, axis=2) final_result = np.squeeze(np.any(result, axis=1, keepdims=True)) print(final_result)
输出结果
[False False False True False False]
浮点数匹配注意事项
由于浮点数存在精度误差,直接用==判断相等可能出现意外错误。建议使用np.allclose()替代==,设置合理的容差:
# 带精度容差的完全匹配判断 result = np.all(np.allclose(array4[:, None], array3, atol=1e-6), axis=2) final_result = np.squeeze(np.any(result, axis=1, keepdims=True))
内容的提问来源于stack exchange,提问作者Lihka_nonem
相关产品推荐
相关产品推荐

