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

1# SPDX-License-Identifier: Apache-2.0 

2"""Utilities for constructing, combining, and visualizing interval trees.""" 

3 

4import collections 

5import heapq 

6import operator 

7from collections.abc import Hashable, Sequence 

8from typing import Any, Callable, Iterable 

9 

10from intervaltree import Interval, IntervalTree 

11 

12type IntervalLike[T] = tuple[T, T] | tuple[T, T, Any] 

13 

14 

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. 

21 

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 

26 

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) 

35 

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) 

39 

40 tagged_inputs = (tagged_tree(i, tree) for i, tree in enumerate(inputs)) 

41 events: list[tuple[Any, bool, int]] = [] 

42 endpoints = [] 

43 

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)) 

48 

49 events.append((left, True, i)) 

50 heapq.heappush(endpoints, (right, i)) 

51 

52 while endpoints: 

53 pt, j = heapq.heappop(endpoints) 

54 events.append((pt, False, j)) 

55 

56 if not events: 

57 return target 

58 

59 left = events[0][0] 

60 active = collections.defaultdict(int) 

61 

62 for event in events: 

63 right, rising, i = event 

64 

65 if active and left < right: 

66 target.addi(left, right, tuple(d for d in data if d in active)) 

67 

68 datum = data[i] 

69 

70 if rising: 

71 active[datum] += 1 

72 elif active[datum] > 1: 

73 active[datum] -= 1 

74 else: 

75 del active[datum] 

76 

77 left = right 

78 

79 return target 

80 

81 

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. 

87 

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) 

93 

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 = [] 

98 

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)) 

103 

104 events.append((iv[0], True)) 

105 heapq.heappush(endpoints, iv[1]) 

106 

107 while endpoints: 

108 events.append((heapq.heappop(endpoints), False)) 

109 

110 next_start = None 

111 depth = 0 

112 

113 for event in events: 

114 depth += 1 if event[1] else -1 

115 

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 

123 

124 return target 

125 

126 

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() 

135 

136 if data is None: 

137 data = (None for _ in range(len(inputs))) 

138 

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) 

143 

144 return target 

145 

146 

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]) 

158 

159 if target.overlaps_point(iv[1]): 

160 target.slice(iv[1]) 

161 

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])) 

166 

167 

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 

190 

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}}") 

195 

196 for i in range(width): 

197 ivs = tree.at(start_from + i * incr) 

198 

199 if not ivs: 

200 continue 

201 

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] 

207 

208 if consensus is not None: 

209 if consensus not in chars: 

210 chars[consensus] = chr(char_counter) 

211 char_counter += 1 

212 

213 char = chars[consensus] 

214 else: 

215 char = none_key_char 

216 

217 line[i] = char 

218 

219 print_func("".join(line)) 

220 

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) 

227 

228 left = f"intervals: {count}" 

229 else: 

230 left = "" 

231 

232 right = f"{end_at} \u2501\u2501\u251b" 

233 print_func(f"{left} {right:>{width - len(left) - 1}}") 

234 

235 if legend_show_keys and chars: 

236 import textwrap 

237 

238 for line in textwrap.wrap(" ".join(f"{v}={k}" for k, v in chars.items()), width=width): 

239 print_func(line)