如何用numpy.where或其他方法获取长度>1的子数组首元素?
获取长度大于1的子数组第一个元素的解决方案
首先咱们先解决你的问题,再分析你之前尝试的问题出在哪。
方法一:列表推导式(最简洁直观)
你的b是包含numpy数组的普通列表,直接遍历每个子数组,判断长度是否大于1,符合条件的就取第一个元素:
import numpy as np # 你的原始数组列表 b = [np.array([0]), np.array([47, 48, 49, 50]), np.array([97, 98, 99])] # 筛选并提取目标元素 result = [arr[0] for arr in b if len(arr) > 1] print(result) # 输出: [47, 97]
方法二:结合numpy.where使用
如果你一定要用numpy.where,需要先获取每个子数组的长度,再生成筛选掩码,最后提取对应元素:
import numpy as np b = [np.array([0]), np.array([47, 48, 49, 50]), np.array([97, 98, 99])] # 生成每个子数组的长度数组 subarray_lengths = np.array([len(arr) for arr in b]) # 生成长度大于1的掩码 mask = subarray_lengths > 1 # 获取符合条件的索引,再提取对应子数组的第一个元素 result = np.array([b[idx][0] for idx in np.where(mask)[0]]) print(result) # 输出: array([47, 97])
分析你之前的错误
你尝试的numpy.where(numpy.array(b).shape > 1)之所以得到错误结果,是因为:
当你把列表b转成numpy数组时,由于子数组长度不一致,得到的是一个dtype为object的一维数组,它的shape是(3,)。你用这个shape和1比较,本质是把元组(3,)和整数1对比,这显然不是你想要的——你需要检查的是每个子数组自己的长度,而不是整个大数组的shape。
内容的提问来源于stack exchange,提问作者user13107
相关产品推荐
相关产品推荐

