numpy.where 异常行为疑问:为何执行条件不满足的分支?
numpy.where 为何会执行条件不满足的分支?
numpy.where的执行逻辑是先完整计算所有分支的表达式,再根据条件筛选结果,而非仅计算符合条件的分支。
以你测试的代码np.where(1<=0, np.sqrt(-1), 0)为例,Python会先执行np.sqrt(-1),这一步直接触发无效值警告,之后才判断条件1<=0不成立,最终返回第二个分支的0.。回到你的实际代码
np.where(x<= 1, np.sqrt(1-x**2), 1):哪怕你认为x的元素都满足x<=1,numpy仍会提前计算整个np.sqrt(1-x**2)数组。如果x中存在元素让1-x**2为负(即x>1),就会触发警告——这其实是在提示你,x数组可能存在不符合你预期的元素,和之前的逻辑假设存在矛盾。对比torch.where:PyTorch的where采用惰性计算,仅对满足条件的位置计算对应分支的表达式,不满足条件的分支不会被执行,因此不会出现这类“无意义”的警告。
内容的提问来源于stack exchange,提问作者math_guy
相关产品推荐
相关产品推荐

