numpy.sqrt带where条件处理整数数组返回异常,是Bug还是用法误解?
这不是Bug,是对
where参数的用法误解 numpy里像np.sqrt这种通用函数(ufunc)的where参数,核心逻辑是:只在where条件为True的位置执行计算,并把结果写入输出数组;条件为False的位置,输出数组的对应元素不会被修改。但你忽略了一个关键:当你没显式指定out参数时,numpy会自动创建一个输出数组,而这个数组的初始值规则和输入类型直接相关:
- 浮点数数组输入:你的示例里输出数组刚好初始化成了原数组的副本(这是特定环境/版本的表现,并非所有情况都这样),所以前两个不满足
a>5的元素保留了原值,第三个元素被替换为sqrt(9)的结果,得到[1.,4.,3.]。 - 整数数组输入:因为
sqrt返回浮点数,numpy会新建一个浮点数数组,但这个数组的初始值是未初始化的(也就是内存里的随机值,你的示例中是0.和0.5),只有第三个满足条件的元素被写入3.,所以出现了奇怪的结果。
如果想实现「只对满足条件的元素计算平方根,其他元素保持原数组值」的需求,正确的做法是显式指定out参数为原数组的副本(注意整数数组要先转成浮点数,因为sqrt返回浮点数):
import numpy as np # 整数数组的正确写法 a = np.array([1, 4, 9]) out = a.astype(float) np.sqrt(a, where=(a>5), out=out) # 输出:array([1., 4., 3.]) # 浮点数数组的可靠写法(显式指定out避免依赖默认行为) a = np.array([1., 4., 9.]) out = a.copy() np.sqrt(a, where=(a>5), out=out) # 输出:array([1., 4., 3.])
内容的提问来源于stack exchange,提问作者Antoine101
相关产品推荐
相关产品推荐

