迭代二维数组时如何使用np.where筛选末尾元素大于1的子数组
问题解决方法
错误根源
你当前逻辑出错的核心原因是对np.where的返回值理解有误:np.where(condition)返回的是满足条件的元素索引组成的元组,只要condition成立,返回的元组内就包含非空的索引数组,Python中非空对象在if判断中都会被视为True,因此无论实际判断结果如何,你的代码都会执行追加操作。
你当前场景下prevalence[-1]是单个数值,prevalence[-1] > 1本身就是布尔标量,完全不需要调用np.where,直接用该表达式做判断即可。
修改方案
直接替换倒数第二行代码为以下内容即可:
if prevalence[-1] > 1:
如果一定要使用np.where实现该逻辑,需要判断返回的索引数组是否非空,写法如下:
if np.where(prevalence[-1] > 1)[0].size > 0:
补充优化
你也可以不用循环逐个判断,先合并所有数组为二维数组后一次性过滤,代码更简洁高效:
import numpy as np import glob # 注意原代码这里缺少右括号,已补全 reps = sorted(glob.glob('C:/Users/Repetitions/*')) all_reps = np.array([np.loadtxt(r, delimiter=',') for r in reps]) # 直接过滤最后一列大于1的所有行 outbreaks = all_reps[all_reps[:, -1] > 1]
内容的提问来源于stack exchange,提问作者Annette
相关产品推荐
相关产品推荐

