You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

高效实现:从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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.13 11:25:03