寻找LeetCode非重叠区间问题解法的错误边缘案例
寻找「无重叠区间」解法的错误边缘案例
我正在解决LeetCode上的「Non-overlapping Intervals」问题,该问题要求计算需要删除的最少区间数量,以得到一组无重叠的区间(需删除的数量即为所求结果)。
我的解决方案步骤如下:
- 基于所有区间构建增强区间树,时间复杂度为O(n log n)
- 遍历每个区间,统计其与其他区间的相交次数(统计结果包含自身相交,因此需减1作为相对指标),这一步时间复杂度同样为O(n log n)
- 将所有区间按该相交次数指标降序排序
- 从排序后的列表中逐个取出区间,使用另一棵区间树显式检查是否重叠,构建无重叠集合,将重叠的区间加入待删除集合
以下是完整的解决方案代码:
from typing import List, Iterable class Interval: def __init__(self, lo: int, hi: int): self.lo = lo self.hi = hi class Node: def __init__(self, interval: Interval, left: 'Node' = None, right: 'Node' = None): self.left = left self.right = right self.interval = interval self.max_hi = interval.hi class IntervalTree: def __init__(self): self.root = None def __add(self, interval: Interval, node:Node) -> Node: if node is None: node = Node(interval) node.max_hi = interval.hi return node if node.interval.lo > interval.lo: node.left = self.__add(interval, node.left) else: node.right = self.__add(interval, node.right) node.max_hi = max(node.left.max_hi if node.left else 0, node.right.max_hi if node.right else 0, node.interval.hi) return node def add(self, lo: int, hi: int): interval = Interval(lo, hi) self.root = self.__add(interval, self.root) def __is_intersect(self, interval: Interval, node: Node) -> bool: if node is None: return False if not (node.interval.lo >= interval.hi or node.interval.hi <= interval.lo): return True if node.left and node.left.max_hi > interval.lo: return self.__is_intersect(interval, node.left) return self.__is_intersect(interval, node.right) def is_intersect(self, lo: int, hi: int) -> bool: interval = Interval(lo, hi) return self.__is_intersect(interval, self.root) def __all_intersect(self, interval: Interval, node: Node) -> Iterable[Interval]: if node is None: yield from () else: if not (node.interval.lo >= interval.hi or node.interval.hi <= interval.lo): yield node.interval if node.left and node.left.max_hi > interval.lo: yield from self.__all_intersect(interval, node.left) yield from self.__all_intersect(interval, node.right) def all_intersect(self, lo: int, hi: int) -> Iterable[Interval]: interval = Interval(lo, hi) yield from self.__all_intersect(interval, self.root) class Solution: def eraseOverlapIntervals(self, intervals: List[List[int]]) -> int: ranged_intervals = [] interval_tree = IntervalTree() for interval in intervals: interval_tree.add(interval[0], interval[1]) for interval in intervals: c = interval_tree.all_intersect(interval[0], interval[1]) ranged_intervals.append((len(list(c))-1, interval)) # 减去自身相交的计数 interval_tree = IntervalTree() res = [] ranged_intervals.sort(key=lambda t: t[0], reverse=True) while ranged_intervals: _, interval = ranged_intervals.pop() if not interval_tree.is_intersect(interval[0], interval[1]): interval_tree.add(interval[0], interval[1]) else: res.append(interval) return len(res)
该解法运行速度足够快,但有时会出现结果差1的错误,比如预期结果为810,而我的结果为811。即使我知道该问题的其他解法,仍希望找到导致该解法失败的边缘案例,若有人能发现此类案例,将不胜感激!
提前感谢任何建设性的意见和思路!
内容的提问来源于stack exchange,提问作者RomanGirin
相关产品推荐
相关产品推荐

