Skip to main content

competitive/data_structure/
persistent_segment_tree.rs

1use 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}