Numpy相等运算符重载与向量化:链式比较真值歧义错误排查
Numpy链式相等运算报错的原因解析
先看你给出的代码示例:
import numpy as np x = np.array([True, False]) (x == False) == True # 正常运行,输出array([False, True]) x == (False == True) # 正常运行,输出array([False, True]) x == False == True # 抛出错误
报错信息:
ValueError: The truth value of an array with more than one element is ambiguous. Use a.any() or a.all()
核心原因:Python链式比较的逻辑和你预期的不一样
你可能以为x == False == True等价于(x == False) == True或者x == (False == True),但Python的链式比较语法规则是:a == b == c等价于(a == b) and (b == c),而不是连续的相等判断。
拆解三种情况的执行逻辑:
(x == False) == True:
先执行x == False得到布尔数组array([False, True]),再将这个数组和True做广播比较,逐个元素判断是否等于True,最终得到结果数组。x == (False == True):
先执行False == True得到单个布尔值False,再让数组x和False逐个元素比较,得到结果数组。x == False == True:
按照Python规则等价于(x == False) and (False == True)。前半部分得到布尔数组array([False, True]),后半部分得到单个布尔值False。Numpy不允许直接用and连接布尔数组和标量——因为and需要单个布尔值作为判断依据,而数组有多个元素,Numpy无法确定你是要判断数组所有元素为真、还是至少一个为真,因此抛出歧义错误,提示你用any()或all()明确判断逻辑。
内容的提问来源于stack exchange,提问作者fblthp
相关产品推荐
相关产品推荐

