在NumPy数组中查找重复序列的匹配索引(含跨数组查询)
解决NumPy重复序列分组与跨数组查询问题
我来帮你搞定这个需求,咱们一步步来实现:
第一步:从源数组提取唯一重复序列并分组索引
你已经有了consec_repeat_starts(a, n)函数来获取重复序列的起始索引,接下来只需要把这些索引按对应的序列分组就行。这里用字典来做分组会很高效,因为序列转成元组后可以作为字典的键:
import numpy as np def group_repeat_sequences(a, n): # 获取所有符合长度要求的重复序列起始索引 starts = consec_repeat_starts(a, n) seq_groups = {} for start in starts: # 把序列转成元组作为字典键(列表不能当键) current_seq = tuple(a[start:start+n]) if current_seq not in seq_groups: seq_groups[current_seq] = [] seq_groups[current_seq].append(start) # 转换成你需要的输出格式:(序列列表, 索引列表)的元组组成的列表 return [(list(seq), indices) for seq, indices in seq_groups.items()]
测试一下第一个数组:
a = np.array([0, 1, 2, 2, 3, 4, 2, 2, 5, 5, 6, 5, 5, 2, 2]) n = 2 # 指定重复序列长度 result1 = group_repeat_sequences(a, n) print(result1) # 输出正好符合你的示例:[([2, 2], [2, 6, 13]), ([5, 5], [8, 11])]
第二步:跨数组查询指定序列的匹配索引
接下来要拿第一步得到的唯一序列,去第二个数组里找所有匹配的起始位置。这里用NumPy的滑动窗口来实现快速匹配,效率很高:
def find_sequence_matches(b, target_seq): seq_len = len(target_seq) if seq_len > len(b): return [] # 序列比数组长,不可能匹配 # 生成数组的滑动窗口,每个窗口长度等于目标序列长度 sliding_windows = np.lib.stride_tricks.sliding_window_view(b, seq_len) # 逐窗口比较是否和目标序列完全一致 match_mask = np.all(sliding_windows == target_seq, axis=1) # 提取所有匹配的起始索引并转成列表 return np.where(match_mask)[0].tolist() def query_sequences_in_array(unique_seqs, b): # unique_seqs是第一步得到的分组结果 query_result = [] for seq, _ in unique_seqs: # 转成元组方便后续比较 seq_tuple = tuple(seq) matches = find_sequence_matches(b, seq_tuple) query_result.append((seq, matches)) return query_result
测试第二个数组:
b = np.array([2, 2, 5, 5, 1, 4, 9, 2, 5, 5, 0, 2, 2, 2]) result2 = query_sequences_in_array(result1, b) print(result2) # 输出符合示例:[([2, 2], [0, 11, 12]), ([5, 5], [2, 8])]
补充:如果还没实现consec_repeat_starts函数
要是你之前的consec_repeat_starts还没写好,这里给你一个高效的实现,支持任意指定的重复序列长度n:
def consec_repeat_starts(a, n): if n == 1: return list(range(len(a))) # 生成相邻元素相等的掩码 equal_adjacent = np.concatenate([[False], a[1:] == a[:-1]]) # 把连续相等的元素分成不同组 group_ids = np.cumsum(~equal_adjacent) # 统计每个组的元素数量 group_counts = np.bincount(group_ids) # 遍历每个组,提取符合长度要求的起始索引 starts = [] current_pos = 0 for count in group_counts: if count >= n: # 一个长度为count的连续组,可以生成count - n + 1个起始索引 starts.extend(range(current_pos, current_pos + count - n + 1)) current_pos += count return starts
这个函数能准确找到所有连续至少n个相同元素的起始位置,完全适配你的需求。
内容的提问来源于stack exchange,提问作者slaw
相关产品推荐
相关产品推荐

