Coverage for core / src / sensorkit / common / interval.py: 67%
133 statements
« prev ^ index » next coverage.py v7.13.5, created at 2026-09-02 00:03 +0000
« prev ^ index » next coverage.py v7.13.5, created at 2026-09-02 00:03 +0000
1# SPDX-License-Identifier: Apache-2.0
2"""Utilities for constructing, combining, and visualizing interval trees."""
4import collections
5import heapq
6import operator
7from collections.abc import Hashable, Sequence
8from typing import Any, Callable, Iterable
10from intervaltree import Interval, IntervalTree
12type IntervalLike[T] = tuple[T, T] | tuple[T, T, Any]
15def combine_interval_trees(
16 *inputs: Iterable[IntervalLike],
17 data: Iterable[Any] | None = None,
18 target: IntervalTree | None = None,
19):
20 """Construct an IntervalTree combining the data fields of all input interval collections.
22 Each collection of intervals must be disjoint. Assuming the inputs are IntervalTrees, this can
23 be achieved using the `merge_overlaps()` method.
24 """
25 target = IntervalTree() if target is None else target
27 match data:
28 case None:
29 data = tuple(None for _ in range(len(inputs)))
30 case Sequence():
31 pass
32 case _:
33 # Make sure the data iterable is materialized for random access.
34 data = tuple(data)
36 # Line sweep, making sure to preserve the input ordering. Each input tree is sorted here.
37 def tagged_tree(i, tree):
38 return sorted((iv[0], i, iv[1]) for iv in tree)
40 tagged_inputs = (tagged_tree(i, tree) for i, tree in enumerate(inputs))
41 events: list[tuple[Any, bool, int]] = []
42 endpoints = []
44 for left, i, right in heapq.merge(*tagged_inputs):
45 while endpoints and endpoints[0][0] <= left:
46 pt, j = heapq.heappop(endpoints)
47 events.append((pt, False, j))
49 events.append((left, True, i))
50 heapq.heappush(endpoints, (right, i))
52 while endpoints:
53 pt, j = heapq.heappop(endpoints)
54 events.append((pt, False, j))
56 if not events:
57 return target
59 left = events[0][0]
60 active = collections.defaultdict(int)
62 for event in events:
63 right, rising, i = event
65 if active and left < right:
66 target.addi(left, right, tuple(d for d in data if d in active))
68 datum = data[i]
70 if rising:
71 active[datum] += 1
72 elif active[datum] > 1:
73 active[datum] -= 1
74 else:
75 del active[datum]
77 left = right
79 return target
82def intersect_interval_trees(
83 *inputs: Iterable[IntervalLike],
84 target: IntervalTree | None = None,
85):
86 """Construct an IntervalTree containing the intersection of all input interval collections.
88 Each collection of intervals must be disjoint. Assuming the inputs are IntervalTrees, this can
89 be achieved using the `merge_overlaps()` method.
90 """
91 target = IntervalTree() if target is None else target
92 interval_key = operator.itemgetter(0, 1)
94 # Line sweep. Each input tree is sorted here.
95 sorted_inputs = (sorted(inp, key=interval_key) for inp in inputs)
96 events: list[tuple[Any, bool]] = []
97 endpoints = []
99 for iv in heapq.merge(*sorted_inputs, key=interval_key):
100 while endpoints and endpoints[0] <= iv[0]:
101 end = heapq.heappop(endpoints)
102 events.append((end, False))
104 events.append((iv[0], True))
105 heapq.heappush(endpoints, iv[1])
107 while endpoints:
108 events.append((heapq.heappop(endpoints), False))
110 next_start = None
111 depth = 0
113 for event in events:
114 depth += 1 if event[1] else -1
116 if depth == len(inputs):
117 next_start = event[0] if next_start is None else next_start
118 elif next_start is not None and next_start < event[0]:
119 target.addi(next_start, event[0])
120 next_start = None
121 elif next_start is not None:
122 next_start = None
124 return target
127def stack_interval_trees(
128 *inputs: Iterable[IntervalLike],
129 data: Iterable[Any] | None = None,
130 target: IntervalTree | None = None,
131):
132 """Add each input interval collection to an IntervalTree, splitting existing intervals."""
133 if target is None:
134 target = IntervalTree()
136 if data is None:
137 data = (None for _ in range(len(inputs)))
139 for iv_data, tree in zip(data, inputs, strict=True):
140 for iv in tree:
141 target.chop(iv[0], iv[1])
142 target.addi(iv[0], iv[1], iv_data)
144 return target
147def stamp_interval_trees(
148 target: IntervalTree,
149 *,
150 stamp: Iterable[IntervalLike],
151 merge_func: Callable[[Any, Any], Any],
152):
153 """Slice the target IntervalTree on "stamp" interval boundaries and merge overlapping data."""
154 for iv in stamp:
155 # Split overlaps at both endpoints so inner pieces are fully enveloped.
156 if target.overlaps_point(iv[0]):
157 target.slice(iv[0])
159 if target.overlaps_point(iv[1]):
160 target.slice(iv[1])
162 # Merge data for all intervals fully inside the stamp range.
163 for enveloped in target.envelop(iv[0], iv[1]):
164 target.remove(enveloped)
165 target.addi(enveloped.begin, enveloped.end, merge_func(enveloped.data, iv[2]))
168def print_interval_tree(
169 tree: IntervalTree,
170 *,
171 width: int = 80,
172 start_from: Any = None,
173 end_at: Any = None,
174 key_func: Callable[[Interval], Hashable] = lambda iv: iv.data,
175 print_func: Callable[..., Any] = print,
176 no_interval_char: str = "_",
177 none_key_char: str = "^",
178 legend: bool = True,
179 legend_show_count: bool = False,
180 legend_show_keys: bool = False,
181):
182 """Render an IntervalTree as a fixed-width ASCII timeline with an optional legend."""
183 line = [no_interval_char] * width
184 char_counter = ord("A")
185 chars = {}
186 start_from = tree.begin() if start_from is None else start_from
187 end_at = tree.end() if end_at is None else end_at
188 span = end_at - start_from
189 incr = span / width
191 if legend:
192 left = f"\u250f\u2501\u2501 {start_from}"
193 right = f"span: {span}"
194 print_func(f"{left} {right:>{width - len(left) - 1}}")
196 for i in range(width):
197 ivs = tree.at(start_from + i * incr)
199 if not ivs:
200 continue
202 if len(ivs) == 1:
203 consensus = key_func(next(iter(ivs)))
204 else:
205 counts = collections.Counter(key_func(iv) for iv in ivs)
206 consensus = counts.most_common(1)[0][0]
208 if consensus is not None:
209 if consensus not in chars:
210 chars[consensus] = chr(char_counter)
211 char_counter += 1
213 char = chars[consensus]
214 else:
215 char = none_key_char
217 line[i] = char
219 print_func("".join(line))
221 if legend:
222 if legend_show_count:
223 ivs = tree.envelop(start_from, end_at)
224 count = len(ivs)
225 count += len(tree.at(start_from) - ivs)
226 count += len(tree.at(end_at) - ivs)
228 left = f"intervals: {count}"
229 else:
230 left = ""
232 right = f"{end_at} \u2501\u2501\u251b"
233 print_func(f"{left} {right:>{width - len(left) - 1}}")
235 if legend_show_keys and chars:
236 import textwrap
238 for line in textwrap.wrap(" ".join(f"{v}={k}" for k, v in chars.items()), width=width):
239 print_func(line)