如何用Numpy简洁地道地获取二维数组的逐行相等结果?
实现Numpy二维数组逐行相等判断的简洁方法
你需要的其实是在逐元素比较的基础上,沿着行的维度(axis=1)判断整行所有元素是否都相等,用all()方法指定axis参数就可以完美解决这个问题。
两种简洁实现方式:
- 先做逐元素比较,再调用数组的
all()方法指定axis=1:
import numpy as np a = np.array([[1, 2], [3, 4], [5, 6]]) b = np.array([[5, 6], [3, 4], [1, 2]]) result = (a == b).all(axis=1) print(result) # 输出:array([False, True, False])
- 用
np.all()函数封装,效果完全一致:
result = np.all(a == b, axis=1) print(result) # 输出同样符合预期的结果
为什么原来的方法不符合预期?
你之前用a == b或者np.equal(a, b)得到的是逐元素的布尔数组,这是因为Numpy的广播机制会默认对每个元素单独比较。而加上all(axis=1)之后,会沿着每行的方向(也就是第二个维度)检查所有元素是否都为True,最终压缩成一维数组,正好对应每行是否完全相等的判断结果。
补充:处理浮点数数组的情况
如果你的数组是浮点数类型,直接用==可能会因为精度问题导致错误判断,这时候可以用np.allclose()来做行级的近似相等判断:
a_float = np.array([[1.0, 2.0], [3.0, 4.0]]) b_float = np.array([[1.0000001, 2.0], [3.0, 4.0]]) result = np.allclose(a_float, b_float, axis=1) print(result) # 输出:array([ True, True])
内容的提问来源于stack exchange,提问作者gojomo
相关产品推荐
相关产品推荐

