competitive/data_structure/
persistent_segment_tree.rs1use super::{Allocator, MemoryPool, Monoid, RangeBoundsExt};
2use std::{
3 fmt::{self, Debug, Formatter},
4 ops::{Range, RangeBounds},
5 ptr::NonNull,
6};
7
8type NodePtr<T> = Option<NonNull<Node<T>>>;
9
10struct Node<T> {
11 children: [NodePtr<T>; 2],
12 value: T,
13}
14
15impl<T> Node<T> {
16 fn new(children: [NodePtr<T>; 2], value: T) -> Self {
17 Self { children, value }
18 }
19}
20
21#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
22#[must_use]
23pub struct PersistentSegmentTreeVersion(usize);
24
25impl PersistentSegmentTreeVersion {
26 fn base() -> Self {
27 Self(0)
28 }
29
30 fn new(version_id: usize) -> Self {
31 Self(version_id)
32 }
33
34 fn index(self) -> usize {
35 self.0
36 }
37}
38
39pub struct PersistentSegmentTree<M>
40where
41 M: Monoid,
42{
43 len: usize,
44 version_roots: Vec<NodePtr<M::T>>,
45 allocator: MemoryPool<Node<M::T>>,
46}
47
48impl<M> Debug for PersistentSegmentTree<M>
49where
50 M: Monoid,
51{
52 fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
53 f.debug_struct("PersistentSegmentTree")
54 .field("len", &self.len)
55 .field("versions", &self.version_roots.len())
56 .finish()
57 }
58}
59
60impl<M> PersistentSegmentTree<M>
61where
62 M: Monoid,
63{
64 #[must_use]
65 pub fn new(len: usize) -> Self {
66 Self {
67 len,
68 version_roots: vec![None],
69 allocator: MemoryPool::new(),
70 }
71 }
72
73 pub fn base(&self) -> PersistentSegmentTreeVersion {
74 PersistentSegmentTreeVersion::base()
75 }
76
77 pub fn len(&self) -> usize {
78 self.len
79 }
80
81 pub fn is_empty(&self) -> bool {
82 self.len == 0
83 }
84
85 fn version_root(&self, version: PersistentSegmentTreeVersion) -> NodePtr<M::T> {
86 *self
87 .version_roots
88 .get(version.index())
89 .expect("invalid version")
90 }
91
92 fn push_version_root(&mut self, root: NodePtr<M::T>) -> PersistentSegmentTreeVersion {
93 let version_id = self.version_roots.len();
94 self.version_roots.push(root);
95 PersistentSegmentTreeVersion::new(version_id)
96 }
97
98 fn allocate_node(&mut self, children: [NodePtr<M::T>; 2], value: M::T) -> NonNull<Node<M::T>> {
99 self.allocator.allocate(Node::new(children, value))
100 }
101
102 fn build_dfs(&mut self, start: usize, end: usize, values: &[M::T]) -> NodePtr<M::T> {
103 if end - start == 1 {
104 return self.leaf_node(values[start].clone());
105 }
106 let mid = (start + end) / 2;
107 let left = self.build_dfs(start, mid, values);
108 let right = self.build_dfs(mid, end, values);
109 self.merge_nodes(left, right)
110 }
111
112 fn leaf_node(&mut self, value: M::T) -> NodePtr<M::T> {
113 Some(self.allocate_node([None, None], value))
114 }
115
116 fn merge_nodes(&mut self, left: NodePtr<M::T>, right: NodePtr<M::T>) -> NodePtr<M::T> {
117 if left.is_none() && right.is_none() {
118 None
119 } else {
120 let value = M::operate(&Self::subtree_value(left), &Self::subtree_value(right));
121 Some(self.allocate_node([left, right], value))
122 }
123 }
124
125 fn subtree_value(node: NodePtr<M::T>) -> M::T {
126 node.map(|node| unsafe { node.as_ref().value.clone() })
127 .unwrap_or_else(M::unit)
128 }
129
130 fn children(node: NodePtr<M::T>) -> [NodePtr<M::T>; 2] {
131 node.map(|node| unsafe { node.as_ref().children })
132 .unwrap_or([None, None])
133 }
134
135 fn point_get_dfs(node: NodePtr<M::T>, start: usize, end: usize, index: usize) -> M::T {
136 let Some(node) = node else {
137 return M::unit();
138 };
139 let node = unsafe { node.as_ref() };
140 if end - start == 1 {
141 node.value.clone()
142 } else {
143 let mid = (start + end) / 2;
144 if index < mid {
145 Self::point_get_dfs(node.children[0], start, mid, index)
146 } else {
147 Self::point_get_dfs(node.children[1], mid, end, index)
148 }
149 }
150 }
151
152 fn fold_dfs(node: NodePtr<M::T>, start: usize, end: usize, range: &Range<usize>) -> M::T {
153 if range.end <= start || end <= range.start {
154 return M::unit();
155 }
156 let Some(node) = node else {
157 return M::unit();
158 };
159 let node = unsafe { node.as_ref() };
160 if range.start <= start && end <= range.end {
161 node.value.clone()
162 } else {
163 let mid = (start + end) / 2;
164 if range.end <= mid {
165 return Self::fold_dfs(node.children[0], start, mid, range);
166 }
167 if mid <= range.start {
168 return Self::fold_dfs(node.children[1], mid, end, range);
169 }
170 let left = Self::fold_dfs(node.children[0], start, mid, range);
171 let right = Self::fold_dfs(node.children[1], mid, end, range);
172 M::operate(&left, &right)
173 }
174 }
175
176 fn partition_point_dfs<P>(
177 node: NodePtr<M::T>,
178 start: usize,
179 end: usize,
180 left: usize,
181 acc: &mut M::T,
182 pred: &mut P,
183 ) -> Option<usize>
184 where
185 P: FnMut(&M::T) -> bool,
186 {
187 if end <= left {
188 return None;
189 }
190 if left <= start {
191 let nacc = M::operate(acc, &Self::subtree_value(node));
192 if pred(&nacc) {
193 *acc = nacc;
194 return None;
195 }
196 if end - start == 1 {
197 return Some(start);
198 }
199 }
200 let mid = (start + end) / 2;
201 let [l, r] = Self::children(node);
202 if let Some(pos) = Self::partition_point_dfs(l, start, mid, left, acc, pred) {
203 Some(pos)
204 } else {
205 Self::partition_point_dfs(r, mid, end, left, acc, pred)
206 }
207 }
208
209 fn rpartition_point_dfs<P>(
210 node: NodePtr<M::T>,
211 start: usize,
212 end: usize,
213 right: usize,
214 acc: &mut M::T,
215 pred: &mut P,
216 ) -> Option<usize>
217 where
218 P: FnMut(&M::T) -> bool,
219 {
220 if right <= start {
221 return None;
222 }
223 if end <= right {
224 let nacc = M::operate(&Self::subtree_value(node), acc);
225 if pred(&nacc) {
226 *acc = nacc;
227 return None;
228 }
229 if end - start == 1 {
230 return Some(end);
231 }
232 }
233 let mid = (start + end) / 2;
234 let [l, r] = Self::children(node);
235 if let Some(pos) = Self::rpartition_point_dfs(r, mid, end, right, acc, pred) {
236 Some(pos)
237 } else {
238 Self::rpartition_point_dfs(l, start, mid, right, acc, pred)
239 }
240 }
241
242 fn set_dfs(
243 &mut self,
244 node: NodePtr<M::T>,
245 start: usize,
246 end: usize,
247 index: usize,
248 value: &M::T,
249 ) -> NodePtr<M::T> {
250 if end - start == 1 {
251 return self.leaf_node(value.clone());
252 }
253 let mid = (start + end) / 2;
254 let mut children = Self::children(node);
255 if index < mid {
256 children[0] = self.set_dfs(children[0], start, mid, index, value);
257 } else {
258 children[1] = self.set_dfs(children[1], mid, end, index, value);
259 }
260 self.merge_nodes(children[0], children[1])
261 }
262
263 fn update_dfs(
264 &mut self,
265 node: NodePtr<M::T>,
266 start: usize,
267 end: usize,
268 index: usize,
269 value: &M::T,
270 ) -> NodePtr<M::T> {
271 if end - start == 1 {
272 return self.leaf_node(M::operate(&Self::subtree_value(node), value));
273 }
274 let mid = (start + end) / 2;
275 let mut children = Self::children(node);
276 if index < mid {
277 children[0] = self.update_dfs(children[0], start, mid, index, value);
278 } else {
279 children[1] = self.update_dfs(children[1], mid, end, index, value);
280 }
281 self.merge_nodes(children[0], children[1])
282 }
283
284 pub fn from_vec(&mut self, v: Vec<M::T>) -> PersistentSegmentTreeVersion {
285 assert_eq!(v.len(), self.len);
286 let root = if self.len == 0 {
287 None
288 } else {
289 self.build_dfs(0, self.len, &v)
290 };
291 self.push_version_root(root)
292 }
293
294 pub fn set(
295 &mut self,
296 version: PersistentSegmentTreeVersion,
297 index: usize,
298 value: M::T,
299 ) -> PersistentSegmentTreeVersion {
300 assert!(index < self.len);
301 let root = self.set_dfs(self.version_root(version), 0, self.len, index, &value);
302 self.push_version_root(root)
303 }
304
305 pub fn update(
306 &mut self,
307 version: PersistentSegmentTreeVersion,
308 index: usize,
309 value: M::T,
310 ) -> PersistentSegmentTreeVersion {
311 assert!(index < self.len);
312 let root = self.update_dfs(self.version_root(version), 0, self.len, index, &value);
313 self.push_version_root(root)
314 }
315
316 #[must_use]
317 pub fn get(&self, version: PersistentSegmentTreeVersion, index: usize) -> M::T {
318 assert!(index < self.len);
319 Self::point_get_dfs(self.version_root(version), 0, self.len, index)
320 }
321
322 #[must_use]
323 pub fn fold<R>(&self, version: PersistentSegmentTreeVersion, range: R) -> M::T
324 where
325 R: RangeBounds<usize>,
326 {
327 let range = range.to_range_bounded(0, self.len).expect("invalid range");
328 if range.is_empty() {
329 M::unit()
330 } else {
331 Self::fold_dfs(self.version_root(version), 0, self.len, &range)
332 }
333 }
334
335 pub fn partition_point_acc<P>(
336 &self,
337 version: PersistentSegmentTreeVersion,
338 left: usize,
339 mut pred: P,
340 ) -> (usize, M::T)
341 where
342 P: FnMut(&M::T) -> bool,
343 {
344 let root = self.version_root(version);
345 let mut acc = M::unit();
346 let pos = if self.len == 0 {
347 None
348 } else {
349 Self::partition_point_dfs(root, 0, self.len, left, &mut acc, &mut pred)
350 };
351 (pos.unwrap_or(self.len), acc)
352 }
353
354 pub fn rpartition_point_acc<P>(
355 &self,
356 version: PersistentSegmentTreeVersion,
357 right: usize,
358 mut pred: P,
359 ) -> (usize, M::T)
360 where
361 P: FnMut(&M::T) -> bool,
362 {
363 let root = self.version_root(version);
364 let mut acc = M::unit();
365 let pos = if self.len == 0 {
366 None
367 } else {
368 Self::rpartition_point_dfs(root, 0, self.len, right, &mut acc, &mut pred)
369 };
370 (pos.unwrap_or(0), acc)
371 }
372
373 #[must_use]
374 pub fn fold_all(&self, version: PersistentSegmentTreeVersion) -> M::T {
375 Self::subtree_value(self.version_root(version))
376 }
377}
378
379#[cfg(test)]
380mod tests {
381 use super::*;
382 use crate::{
383 algebra::ConcatenateOperation,
384 tools::{WithEmptySegment as Wes, Xorshift},
385 };
386
387 const N: usize = 12;
388 const Q: usize = 2_000;
389 const SIGMA: u8 = 6;
390
391 fn rand_word(rng: &mut Xorshift) -> Vec<u8> {
392 let len = rng.random(0..4usize);
393 (0..len).map(|_| rng.random(0..SIGMA)).collect()
394 }
395
396 #[test]
397 fn test_persistent_segment_tree_random_non_commutative() {
398 let mut rng = Xorshift::default();
399 let mut segtree: PersistentSegmentTree<ConcatenateOperation<u8>> =
400 PersistentSegmentTree::new(N);
401 let initial: Vec<_> = (0..N).map(|_| rand_word(&mut rng)).collect();
402 let mut versions = vec![segtree.base(), segtree.from_vec(initial.clone())];
403 let mut states = vec![vec![Vec::new(); N], initial];
404
405 for _ in 0..Q {
406 let base_version = rng.random(0..versions.len());
407 let index = rng.random(0..N);
408 let mut state = states[base_version].clone();
409
410 if rng.gen_bool(0.5) {
411 let value = rand_word(&mut rng);
412 state[index] = value.clone();
413 versions.push(segtree.set(versions[base_version], index, value));
414 } else {
415 let value = rand_word(&mut rng);
416 state[index].extend_from_slice(&value);
417 versions.push(segtree.update(versions[base_version], index, value));
418 }
419 states.push(state);
420
421 let version = rng.random(0..versions.len());
422 let index = rng.random(0..N);
423 let (start, end) = rng.random(Wes(N));
424 let expected: Vec<_> = states[version][start..end]
425 .iter()
426 .flat_map(|word| word.iter().copied())
427 .collect();
428 let expected_all: Vec<_> = states[version]
429 .iter()
430 .flat_map(|word| word.iter().copied())
431 .collect();
432
433 assert_eq!(
434 segtree.get(versions[version], index),
435 states[version][index]
436 );
437 assert_eq!(segtree.fold(versions[version], start..end), expected);
438 assert_eq!(segtree.fold_all(versions[version]), expected_all);
439
440 let left = rng.random(0..=N);
441 let limit = rng.random(1..=N * 4);
442 let mut expected_acc = Vec::new();
443 let mut expected_pos = left;
444 while expected_pos < N {
445 let mut nacc = expected_acc.clone();
446 nacc.extend_from_slice(&states[version][expected_pos]);
447 if nacc.len() < limit {
448 expected_acc = nacc;
449 expected_pos += 1;
450 } else {
451 break;
452 }
453 }
454 assert_eq!(
455 segtree.partition_point_acc(versions[version], left, |acc| acc.len() < limit),
456 (expected_pos, expected_acc)
457 );
458
459 let right = rng.random(0..=N);
460 let limit = rng.random(1..=N * 4);
461 let mut expected_acc = Vec::new();
462 let mut expected_pos = right;
463 while expected_pos > 0 {
464 let mut nacc = states[version][expected_pos - 1].clone();
465 nacc.extend_from_slice(&expected_acc);
466 if nacc.len() < limit {
467 expected_acc = nacc;
468 expected_pos -= 1;
469 } else {
470 break;
471 }
472 }
473 assert_eq!(
474 segtree.rpartition_point_acc(versions[version], right, |acc| acc.len() < limit),
475 (expected_pos, expected_acc)
476 );
477 }
478 }
479}