Python中NumPy数组传入函数时if条件的正确实现问题
问题描述
定义如下Python函数:
def funtion(x, bb, aa): if x>aa: res = aa else: xxr = x/aa res = bb*(1.5*xxr-0.5*xxr**3) return res
执行以下代码时:
xx = np.linspace(0,49,50) yy = funct(xx,74,33)
出现错误:
The truth value of an array with more than one element is ambiguous. Use a.any() or a.all()
尝试用a.any()或a.all()后结果不符合预期,得到的结果中x>33的部分没有被替换为33;期望结果是所有x>33的元素返回常量33,其余元素按公式计算。
解决方案
普通if语句无法直接处理NumPy数组的逐元素判断,需要使用向量化操作实现需求,以下两种方法均无需循环,适合被其他函数调用:
方法1:使用np.where(推荐)
np.where可对数组的每个元素进行条件判断,返回对应位置的结果:
import numpy as np def function(x, bb, aa): xxr = x / aa # 逐元素判断:x>aa时返回aa,否则计算表达式 res = np.where(x > aa, aa, bb * (1.5 * xxr - 0.5 * xxr**3)) return res
方法2:布尔索引替换
先计算所有元素的表达式结果,再将满足条件的位置替换为aa:
import numpy as np def function(x, bb, aa): xxr = x / aa res = bb * (1.5 * xxr - 0.5 * xxr**3) # 替换x>aa的元素为aa res[x > aa] = aa return res
验证结果
调用修改后的函数:
xx = np.linspace(0,49,50) yy = function(xx,74,33) print(yy)
将得到期望的结果:
array([ 0. , 3.36260678, 6.71903609, 10.06311044, 13.38865236, 16.68948438, 19.959429 , 23.19230876, 26.38194618, 29.52216379, 32.60678409, 35.62962963, 38.58452292, 41.46528647, 44.26574283, 46.9797145 , 49.60102401, 52.12349389, 54.54094666, 56.84720483, 59.03609094, 61.1014275 , 63.03703704, 64.83674208, 66.49436514, 68.00372875, 69.35865542, 70.55296769, 71.58048808, 72.4350391 , 73.11044328, 73.60052314, 73.8991012 , 74. , 33. , 33. , 33. , 33. , 33. , 33. , 33. , 33. , 33. , 33. , 33. , 33. , 33. , 33. , 33. , 33. ])
内容的提问来源于stack exchange,提问作者diedro
相关产品推荐
相关产品推荐

