numpy数组按索引匹配检索元素出错,如何修正?
解决numpy按分组索引元组数组的问题
这个问题我之前也踩过类似的坑,核心是没搞清楚numpy二维数组的索引逻辑和np.take的参数用法!
首先帮你分析问题出在哪:你的my_array其实是一个形状为(6, 2)的二维数组(每个原始元组被拆成了数组的一列),而np.where(my_array_values == v)返回的是一个包含行索引的元组(比如当v=1时,返回(array([0, 3]),))。直接把这个元组传给np.take,它会默认把数组扁平化后取元素,自然就得到了奇怪的输出。
接下来给你两种简单的解决方案:
方法一:直接用布尔索引(最直观推荐)
这是numpy中分组取值最常用的方式,不需要绕弯用np.take和np.where,直接用布尔掩码筛选行即可:
import numpy as np my_array = np.array([('AA','11'),('BB','22'),('CC','33'),('DD','44'),('EE','55'),('FF','66')]) my_array_values = np.array([1,2,3,1,3,2]) my_array_values_unique = np.array([1,2,3]) for v in my_array_values_unique: # 用布尔条件筛选对应行 filtered_rows = my_array[my_array_values == v] # 转成元组列表格式,和预期完全匹配 print([tuple(row) for row in filtered_rows])
运行后输出完全符合你的预期:
[('AA', '11'), ('DD', '44')] [('BB', '22'), ('FF', '66')] [('CC', '33'), ('EE', '55')]
方法二:修正np.take的用法
如果你一定要用np.take实现,需要注意两个点:
- 指定
axis=0,表示按行取值; - 从
np.where返回的元组中取出实际的索引数组(因为np.where返回的是(索引数组,)这样的元组)。
代码如下:
for v in my_array_values_unique: # 取出np.where返回的索引数组 target_indices = np.where(my_array_values == v)[0] # 指定axis=0按行取元素 result = np.take(my_array, target_indices, axis=0) print([tuple(row) for row in result])
这样也能得到和预期一致的输出。
补充:为什么原代码会出错?
原代码中np.take(my_array, np.where(my_array_values == v)),np.where返回的是(array([0,3]),)这样的元组,np.take默认会把数组扁平化(axis=None),然后把元组里的每个值当成扁平化后的索引取元素——扁平化后的数组是['AA','11','BB','22','CC','33','DD','44','EE','55','FF','66'],取索引0和3的元素就是'AA'和'22',这就是你看到奇怪输出的原因。
内容的提问来源于stack exchange,提问作者Nakeuh
相关产品推荐
相关产品推荐

