如何在Python中筛选重叠区间组内的最长原区间?
保留重叠区间组中的最长原始区间
给定元组形式的区间列表:
data = [(5,10), (3,10), (13,15), (12,18), (20,29), (25,30)]
每个元组代表一个区间(起始值、结束值),区间间可能存在重叠。需求是:将所有重叠的区间归为一组,每组仅保留原始区间中长度最长的那个,不能合并区间(比如(20,30)是无效输出,因为它不在原始列表中)。测试用例的预期输出为:
output = [(3,10), (12,18), (20,29)]
我曾尝试用NetworkX实现,但该方案扩展性不佳,且不想依赖NetworkX:
import networkx as nx data = [(5,10), (3,10), (13,15), (12,18), (20,29), (25,30)] graph = nx.Graph() n = len(data) for i, a in enumerate(data): a_seq = set(range(a[0], a[1] + 1)) for j in range(i+1, n): b = data[j] b_seq = set(range(b[0], b[1] + 1)) n_overlap = len(a_seq & b_seq) if n_overlap: graph.add_edge(a, b, weight=n_overlap) output = list() for nodes in nx.connected_components(graph): lengths = dict() for node in nodes: start, end = node lengths[node] = end - start longest_interval, length_of_interval = sorted(lengths.items(), key=lambda x:x[1], reverse=True)[0] output.append(longest_interval)
以下是不用NetworkX的更优实现方案,分别基于Python标准库、NumPy和Pandas:
一、Python标准库实现
思路
先按区间起始值排序,再遍历分组重叠区间,每组筛选出长度最长的原始区间。排序时额外按长度降序,能让同起始的长区间优先进入分组,核心通过维护当前组的最大结束值判断重叠。
代码实现
data = [(5,10), (3,10), (13,15), (12,18), (20,29), (25,30)] # 按起始值升序、长度降序排序,便于后续分组 sorted_intervals = sorted(data, key=lambda x: (x[0], -(x[1] - x[0]))) if not sorted_intervals: print([]) else: result = [] current_group = [sorted_intervals[0]] current_max_end = sorted_intervals[0][1] for interval in sorted_intervals[1:]: start, end = interval # 当前区间与组内区间重叠,加入组 if start <= current_max_end: current_group.append(interval) # 更新组内最大结束值,用于后续重叠判断 if end > current_max_end: current_max_end = end else: # 当前组结束,筛选最长区间加入结果 longest = max(current_group, key=lambda x: x[1] - x[0]) result.append(longest) # 初始化新组 current_group = [interval] current_max_end = end # 处理最后一组 longest = max(current_group, key=lambda x: x[1] - x[0]) result.append(longest) print(result) # 输出: [(3, 10), (12, 18), (20, 29)]
二、NumPy实现
思路
将区间转为NumPy数组,通过数组操作标记重叠组ID,再按组筛选出长度最长的原始区间。利用cumsum和maximum.accumulate高效完成分组标记,避免循环遍历。
代码实现
import numpy as np data = [(5,10), (3,10), (13,15), (12,18), (20,29), (25,30)] # 转为数组并保留原始索引 arr = np.array(data) original_indices = np.arange(len(arr)) # 按起始值排序,获取排序后的索引 sorted_idx = np.argsort(arr[:, 0]) sorted_arr = arr[sorted_idx] sorted_original_indices = original_indices[sorted_idx] # 计算每组的最大结束值,标记新组的起始点 max_ends = np.maximum.accumulate(sorted_arr[:, 1]) # 第一个元素为新组,后续元素若起始值大于前一组最大结束值则为新组 new_group_flags = np.concatenate([[True], sorted_arr[1:, 0] > max_ends[:-1]]) group_ids = np.cumsum(new_group_flags) # 遍历每个组,筛选最长原始区间 result = [] for gid in np.unique(group_ids): # 获取当前组的所有元素 group_mask = group_ids == gid group_intervals = sorted_arr[group_mask] # 计算长度并找到最长区间的索引 lengths = group_intervals[:, 1] - group_intervals[:, 0] longest_idx = np.argmax(lengths) # 找到对应的原始区间 original_idx = sorted_original_indices[group_mask][longest_idx] result.append(data[original_idx]) print(result) # 输出: [(3, 10), (12, 18), (20, 29)]
三、Pandas实现
思路
用DataFrame存储区间的起始、结束、长度及原始索引,排序后通过累积最大值和分组标记识别重叠组,最后按组提取最长的原始区间。
代码实现
import pandas as pd data = [(5,10), (3,10), (13,15), (12,18), (20,29), (25,30)] # 构建DataFrame,补充长度和原始索引列 df = pd.DataFrame(data, columns=["start", "end"]) df["length"] = df["end"] - df["start"] df["original_idx"] = df.index # 按起始值排序 df_sorted = df.sort_values("start").reset_index(drop=True) # 计算每组的最大结束值,标记新组 df_sorted["max_end"] = df_sorted["end"].cummax() # 第一个元素为新组,后续元素若起始值大于前一组最大结束值则为新组 df_sorted["is_new_group"] = df_sorted["start"] > df_sorted["max_end"].shift(1, fill_value=-float("inf")) df_sorted["group_id"] = df_sorted["is_new_group"].cumsum() # 按组分组,取每组长度最大的行,提取原始区间 result = ( df_sorted.groupby("group_id") .apply(lambda group: group.loc[group["length"].idxmax()]) .apply(lambda row: (row["start"], row["end"])) .tolist() ) print(result) # 输出: [(3, 10), (12, 18), (20, 29)]
内容的提问来源于stack exchange,提问作者O.rka
相关产品推荐
相关产品推荐

