高效实现:从Numpy整数数组获取各索引对应的位置列表
高效获取Numpy数组中各索引的出现位置
给定形状为d x 1的int型Numpy数组x,数组元素为随机索引,取值范围(0, n)(n为最大可能索引,且n << len(x))。需求是获取每个索引在x中对应的出现位置。
示例代码:
import numpy as np n = 3 x = np.array([2, 0, 3, 3, 2]).reshape(-1, 1) out = required_fn(x, n) # out 应为 [[1], [], [0, 4], [2,3]] # 说明:索引0出现在位置1,索引1无出现,索引2出现在位置0、4,索引3出现在位置2、3
原尝试的x[x == np.array(list(range(n)))]无法实现需求,因此需要更高效的解决方案。
高效实现方案(面向数据加载器的性能优化)
由于功能将用于数据加载器,需优先选择向量化操作避免Python层循环开销,以下是基于Numpy排序与二分查找的高效实现:
import numpy as np def required_fn(x, n): # 扁平化输入数组(处理d x 1的形状) x_flat = x.ravel() # 生成数组元素的位置索引 pos = np.arange(len(x_flat)) # 对x的值排序,同时得到对应的位置索引 sorted_idx = np.argsort(x_flat) sorted_x = x_flat[sorted_idx] sorted_pos = pos[sorted_idx] # 用二分查找确定每个目标索引在排序后数组中的起止位置 targets = np.arange(n + 1) # 覆盖0到n的所有索引 split_points = np.searchsorted(sorted_x, targets) # 分割位置数组得到每个索引对应的出现位置列表 out = [] for i in range(n + 1): start, end = split_points[i], split_points[i + 1] out.append(sorted_pos[start:end].tolist()) return out
方案优势
- 全向量化核心操作:
argsort和searchsorted均为Numpy底层优化的C实现,处理大规模数组时远快于Python循环。 - 时间复杂度更优:整体复杂度为
O(d log d)(排序开销),当n << len(x)时,远优于循环调用np.where的O(n*d)复杂度。 - 覆盖所有索引:自动处理未出现的索引,返回空列表,完全匹配需求。
备选简洁方案(适用于n极小场景)
如果n非常小(如个位数),也可以用更简洁的循环+np.where实现,代码可读性更高:
import numpy as np def required_fn(x, n): x_flat = x.ravel() return [np.where(x_flat == i)[0].tolist() for i in range(n + 1)]
内容的提问来源于stack exchange,提问作者OlorinIstari
相关产品推荐
相关产品推荐

