高效追踪np数组中已选样本索引的最优方法
高效实现无重复随机选取numpy数组行的方法
你的原始实现中,list.remove(i)是**O(n)**时间复杂度的操作——每次删除元素都要遍历列表找到目标位置,当数组规模变大、选取次数增多时,效率会明显下降。下面是两种更优的方案:
方案1:一次性生成所有不重复索引(已知选取次数)
如果提前知道要选多少个样本,直接用numpy的随机数生成器一次性选出所有不重复索引,再逐个处理。这种方法利用numpy的底层优化,效率远高于循环删除列表元素。
方法A:用rng.choice指定replace=False
import numpy as np x = np.arange(100_000).reshape((10_000,10)) rng = np.random.default_rng(seed=42) # 一次性生成1000个不重复的随机索引 selected_indices = rng.choice(len(x), size=1000, replace=False) for i in selected_indices: # 处理x[i],例如: print(x[i].sum())
方法B:用rng.permutation洗牌后取前k个
如果需要遍历所有行的随机顺序,或者k接近数组总行数,用全洗牌更高效:
shuffled_indices = rng.permutation(len(x)) # 取前1000个索引处理 for i in shuffled_indices[:1000]: # 处理x[i] pass
方案2:动态追踪剩余索引(未知选取次数)
如果选取次数不固定(比如中途可能停止选样),可以用Fisher-Yates洗牌的逐步实现——每次随机选一个剩余索引,将其交换到当前剩余数组的末尾,再缩小剩余范围。所有操作都是**O(1)**时间复杂度,完全避免删除开销。
import numpy as np x = np.arange(100_000).reshape((10_000,10)) rng = np.random.default_rng(seed=42) total_rows = len(x) remaining_indices = np.arange(total_rows) current_last = total_rows - 1 for _ in range(1000): # 随机选一个剩余范围内的位置 pick_pos = rng.integers(0, current_last + 1) selected_idx = remaining_indices[pick_pos] # 处理选中的行 # do something with x[selected_idx] # 将选中的索引交换到末尾,缩小剩余范围 remaining_indices[pick_pos], remaining_indices[current_last] = remaining_indices[current_last], remaining_indices[pick_pos] current_last -= 1
方案对比
- 原始方法:每次删除操作O(n),k次选样总时间O(k*n),仅适合极小规模数组。
- 方案1:一次性生成索引,时间复杂度O(n)(permutation)或O(k)(choice优化实现),代码简洁,适合已知选样次数的场景。
- 方案2:每次选样O(1),总时间O(k),适合动态选样、次数不确定的场景,数组规模越大优势越明显。
内容的提问来源于stack exchange,提问作者DeltaIV
相关产品推荐
相关产品推荐

