用NumPy单行代码替代循环实现基于预测值的条件阈值处理
用NumPy单行代码替代阈值处理循环
需求:在预测数组中,根据每个预测值对应的特定阈值进行处理——若预测值为k,则将对应的概率值与th[k]比较,若概率大于阈值则将该位置的预测值置为0,否则保留原预测值。原循环逻辑可行,需要更简洁的NumPy单行实现。
原代码:
import numpy as np y_pred = np.array([1, 2, 2, 1, 1, 3, 3]) y_prob = np.array([0.5, 0.5, 0.75, 0.25, 0.75, 0.60, 0.40]) th = [0, 0.4, 0.7, 0.5] z_true = np.array([0, 2, 0, 1, 0, 0, 3]) z_pred = y_pred.copy() # 需要替代的循环 for i in range(len(z_pred)): if y_prob[i] > th[y_pred[i]]: z_pred[i] = 0 print(z_pred)
解决方案:NumPy单行实现
利用NumPy的索引广播和条件赋值特性,一行完成逻辑:
z_pred = np.where(y_prob > np.array(th)[y_pred], 0, y_pred)
代码解释:
np.array(th)[y_pred]:将阈值列表转为NumPy数组后,通过y_pred的索引直接取出每个预测值对应的阈值,生成和y_prob长度一致的阈值数组。y_prob > np.array(th)[y_pred]:逐元素比较概率值与对应阈值,得到标记需置0位置的布尔数组。np.where(condition, x, y):根据布尔数组,满足条件的位置取x(即0),不满足的取y(即原预测值y_pred),直接生成最终结果数组。
运行结果与原循环完全一致:[0 2 0 1 0 0 3]
内容的提问来源于stack exchange,提问作者Triceratops
相关产品推荐
相关产品推荐

