如何快速获取子列表元素在大列表中的所有位置?
高效获取子序列元素在大序列中的位置(Numpy/Python)
需求:获取子序列int_sequence中所有元素在大序列int_vocabulary中的位置,追求极致计算效率。
初始测试代码:
from datetime import datetime as dt import numpy as np from numba import jit, int64 n, N = 20, 120000 int_vocabulary = np.array(range(N)) np.random.shuffle(int_vocabulary) # 制造非有序场景 int_sequence = np.random.choice(int_vocabulary, n, replace=False)
已尝试四种方法及测试代码:
@jit(int64[:](int64[:], int64[:], int64, int64)) def check(int_sequence, int_vocabulary, n, N): all_indices = np.full(n, N) for xi in range(n): for i in range(N): if int_sequence[xi] == int_vocabulary[i]: all_indices[xi] = i return all_indices t0 = dt.now() for _ in range(10): all_indices0 = check(int_sequence, int_vocabulary, n, N) t0 = (dt.now() - t0).total_seconds() print("numba : ", t0) t0 = dt.now() for _ in range(10): mask = np.full(len(int_vocabulary), False) for x in int_sequence: mask += int_vocabulary == x all_indices1 = np.flatnonzero(mask) t0 = (dt.now() - t0).total_seconds() print("numpy :", t0) t0 = dt.now() for _ in range(10): all_indices2 = np.array([i for i, x in enumerate(int_vocabulary) if x in int_sequence]) t0 = (dt.now() - t0).total_seconds() print("list comprehension : ", t0) t0 = dt.now() for _ in range(10): mask = np.sum(np.array([int_vocabulary == x for x in int_sequence]), axis=0) all_indices3 = np.flatnonzero(mask) t0 = (dt.now() - t0).total_seconds() print("mixed numpy + list comprehension : ", t0) assert np.sum(all_indices0) == np.sum(all_indices1) assert np.sum(all_indices1) == np.sum(all_indices2) assert np.sum(all_indices2) == np.sum(all_indices3)
10次重复测试结果:
numba : 0.028039 numpy : 0.011616 list comprehension : 3.116753 mixed numpy + list comprehension : 0.032301
更快的优化方案
方案1:字典映射(O(N)预处理 + O(n)查询)
利用字典建立元素到索引的映射,预处理一次后查询极快:
t0 = dt.now() for _ in range(10): vocab_map = {val: idx for idx, val in enumerate(int_vocabulary)} all_indices4 = np.array([vocab_map[x] for x in int_sequence]) t0 = (dt.now() - t0).total_seconds() print("dict mapping : ", t0)
该方法在测试场景下耗时通常远低于0.001秒,是目前最快的方案之一,尤其适合需要多次查询的场景。
方案2:Numpy排序+二分查找(O(N log N)预处理 + O(n log N)查询)
通过对大序列排序并记录原始索引,再用二分查找定位子序列元素:
t0 = dt.now() for _ in range(10): sorted_indices = np.argsort(int_vocabulary) sorted_vocab = int_vocabulary[sorted_indices] pos = np.searchsorted(sorted_vocab, int_sequence) all_indices5 = sorted_indices[pos] t0 = (dt.now() - t0).total_seconds() print("numpy sort + searchsorted : ", t0)
该方案耗时通常在0.002-0.005秒左右,纯Numpy实现,无需额外依赖。
方案对比说明
- 字典映射:预处理和查询的时间复杂度最优,适合单次预处理多次查询的场景,代码简洁直观。
- Numpy排序+二分:纯向量运算,避免Python循环开销,适合需要保持Numpy数组操作风格的场景。
内容的提问来源于stack exchange,提问作者FraSchelle
相关产品推荐
相关产品推荐

