如何对object类型NumPy数组逐行查找超阈值的元素索引
报错原因
你当前的NumPy数组是dtype=object类型,数组的每个元素都是Python原生list,并非数值型的NumPy子数组,NumPy不支持直接将list和浮点数做>大小比较,因此触发类型错误。
涉及的数组结构:
array([list([0.28552457, 0.28552457, 0.28552457, 0.28552457, 0.28552457]), list([0.71641791, 0.71641791, 0.71641791, 0.69565217, 0.69565217]), list([0.95626478, 0.95626478, 0.95513577, 0.95513577, 0.95513577]), ..., list([0.14285714, 0.14285714, 0.14285714, 0.14285714, 0.13793103]), list([0.73846154, 0.73846154, 0.73846154, 0.71641791, 0.71641791]), list([0.72727273, 0.72727273, 0.72727273, 0.70588235, 0.70588235])], dtype=object)
原执行代码:
np.argwhere(y>0.5)
报错信息:
TypeError: '>' not supported between instances of 'list' and 'float'
实现方案
根据数组结构选择对应方法即可:
- 所有子列表长度一致时(你的示例属于该场景),先转成标准数值二维数组再执行原逻辑,运行效率最高:
# 转换为float类型二维数值数组 y = np.array(y.tolist(), dtype=np.float64) # 直接执行查找,返回结果每行对应*行索引*和*列索引* z = np.argwhere(y > 0.5)
- 存在子列表长度不一致的情况时,逐行遍历处理:
res = [] for row_id, row in enumerate(y): # 逐行转数组比较,收集符合条件的列索引 for col_id in np.argwhere(np.array(row) > 0.5).flatten(): res.append([row_id, col_id]) z = np.array(res)
内容的提问来源于stack exchange,提问作者sudojarvis
相关产品推荐
相关产品推荐

