如何在numpy.where中计算列值?实现自定义公式及多条件逻辑
实现指引:基于Status列调整计算逻辑
一、修正Project列计算逻辑
按需求,Project列的计算式为(y - x) * factor * Status,其中Status是数组a第三列的原始数值,无需做布尔判断。修正要点:
- 移除多余的
statusTrue布尔判断,直接使用a[:,2]的原始值参与计算 - 利用NumPy广播机制实现维度匹配:
np.diff(arr_xy, axis=1)得到每个(x,y)对的差值(形状为(N,1)),与a[:,2](形状为(M,))自动广播为(N,M)的矩阵,替代低效的np.tile - 调整嵌套
np.where的逻辑:当rangeExists为真时直接计算目标值,否则返回-factor
二、完善OtherCalc列计算
otherCalc的目标式为factor * (1 - x/y),修正维度匹配问题:
- 直接计算
arr_xy[:,0]/arr_xy[:,1]得到每个(x,y)对的比值(形状为(N,)),转置为(N,1)后与factor相乘,再通过广播匹配到(N,M)的形状 - 用广播替代
np.tile,避免重复复制数据导致的内存浪费
三、完整修正代码
import numpy as np # array of all valid permutations of x and y arr_xy = np.array(np.meshgrid(lb.round(2), ub.round(2))).T.reshape(-1, 2) arr_xy = arr_xy[arr_xy[:, 0] < arr_xy[:, 1]] # lbExists - (x >= row[0] and x <= row[1]) lbExists = (arr_xy[:, 0][:, np.newaxis] >= a[:, 0]) * \ (arr_xy[:, 0][:, np.newaxis] <= a[:, 1]) # rangeExists - (x >= row[0] and y <= row[1]) rangeExists = lbExists * (arr_xy[:, 1][:, np.newaxis] <= a[:, 1]) # ------------------- 修正Project列计算 ------------------- # 直接计算(y-x)*factor*Status,利用广播匹配维度 project_valid = np.diff(arr_xy, axis=1) * factor * a[:,2] # 按条件赋值:lbExists为真时,根据rangeExists选择对应值,否则为0 project = np.where(lbExists, np.where(rangeExists, project_valid, -factor), 0) # ------------------- 完善OtherCalc列计算 ------------------- # 计算factor*(1 - x/y),转置后通过广播匹配维度 otherCalc_valid = factor * (1 - arr_xy[:,0]/arr_xy[:,1])[:, np.newaxis] # 按条件赋值 otherCalc = np.where(lbExists, np.where(rangeExists, otherCalc_valid, -factor), 0) # combine variables arr_out = np.hstack([ # permutations of upper and lower bound np.vstack([arr_xy] * a.shape[0]), # repeated values of Min and Max np.repeat(a, arr_xy.shape[0], axis=0), # lbExists 2d -> 1d lbExists.T.reshape(-1)[:, np.newaxis], # rangeExists 2d -> 1d rangeExists.T.reshape(-1)[:, np.newaxis], # project 2d -> 1d project.T.reshape(-1)[:, np.newaxis].round(2), # other calc otherCalc.T.reshape(-1)[:, np.newaxis].round(2)])
内容的提问来源于stack exchange,提问作者Jared King
相关产品推荐
相关产品推荐

