Skip to main content

competitive/data_structure/
dual_segment_tree.rs

1use super::{MonoidAct, RangeBoundsExt, Unital};
2use std::{
3    fmt::{self, Debug, Formatter},
4    mem::replace,
5    ops::RangeBounds,
6};
7
8pub struct DualSegmentTree<M>
9where
10    M: MonoidAct,
11{
12    n: usize,
13    keys: Vec<M::Key>,
14    lazy: Vec<M::Act>,
15}
16
17impl<M> Clone for DualSegmentTree<M>
18where
19    M: MonoidAct<Key: Clone>,
20{
21    fn clone(&self) -> Self {
22        Self {
23            n: self.n,
24            keys: self.keys.clone(),
25            lazy: self.lazy.clone(),
26        }
27    }
28}
29
30impl<M> Debug for DualSegmentTree<M>
31where
32    M: MonoidAct<Key: Debug, Act: Debug>,
33{
34    fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
35        f.debug_struct("DualSegmentTree")
36            .field("n", &self.n)
37            .field("keys", &self.keys)
38            .field("lazy", &self.lazy)
39            .finish()
40    }
41}
42
43impl<M> DualSegmentTree<M>
44where
45    M: MonoidAct<Key: Clone, Act: PartialEq>,
46{
47    pub fn new(len: usize, key: M::Key) -> Self {
48        let n = len.next_power_of_two();
49        Self {
50            n,
51            keys: vec![key; len],
52            lazy: vec![M::unit(); n],
53        }
54    }
55    pub fn from_keys(keys: impl ExactSizeIterator<Item = M::Key>) -> Self {
56        let keys: Vec<_> = keys.collect();
57        let n = keys.len().next_power_of_two();
58        Self {
59            n,
60            keys,
61            lazy: vec![M::unit(); n],
62        }
63    }
64    fn update_at(&mut self, k: usize, a: &M::Act) {
65        if k < self.n {
66            M::operate_assign(&mut self.lazy[k], a);
67        } else {
68            M::act_assign(&mut self.keys[k - self.n], a);
69        }
70    }
71    fn propagate_at(&mut self, k: usize) {
72        let a = replace(&mut self.lazy[k], M::unit());
73        if !M::ActMonoid::is_unit(&a) {
74            self.update_at(2 * k, &a);
75            self.update_at(2 * k + 1, &a);
76        }
77    }
78    pub fn update<R>(&mut self, range: R, a: M::Act)
79    where
80        R: RangeBounds<usize>,
81    {
82        let range = range
83            .to_range_bounded(0, self.keys.len())
84            .expect("invalid range");
85        if range.is_empty() || M::ActMonoid::is_unit(&a) {
86            return;
87        }
88        let mut l = range.start + self.n;
89        let mut r = range.end + self.n;
90        for i in (1..=self.n.trailing_zeros()).rev() {
91            if (l >> i) << i != l {
92                self.propagate_at(l >> i);
93            }
94            if (r >> i) << i != r && ((l >> i) << i == l || l >> i != (r - 1) >> i) {
95                self.propagate_at((r - 1) >> i);
96            }
97        }
98        while l < r {
99            if l & 1 != 0 {
100                self.update_at(l, &a);
101                l += 1;
102            }
103            if r & 1 != 0 {
104                r -= 1;
105                self.update_at(r, &a);
106            }
107            l >>= 1;
108            r >>= 1;
109        }
110    }
111    pub fn get(&self, k: usize) -> M::Key {
112        let mut value = self.keys[k].clone();
113        let mut k = (k + self.n) >> 1;
114        while k > 0 {
115            value = M::act(&value, &self.lazy[k]);
116            k >>= 1;
117        }
118        value
119    }
120    pub fn set(&mut self, k: usize, value: M::Key) {
121        assert!(k < self.keys.len());
122        let index = k + self.n;
123        for i in (1..=self.n.trailing_zeros()).rev() {
124            self.propagate_at(index >> i);
125        }
126        self.keys[k] = value;
127    }
128}
129
130#[cfg(test)]
131mod tests {
132    use super::*;
133    use crate::{
134        algebra::LinearAct,
135        num::mint_basic::MInt998244353 as M,
136        tools::{
137            Xorshift,
138            testutil::{exhaustive_sequences, sample_usize},
139        },
140    };
141
142    #[test]
143    fn test_dual_segment_tree() {
144        let mut rng = Xorshift::default();
145        for n in sample_usize(&mut rng, 5, 0..=65, 10) {
146            let updates: Vec<_> = (0..=n)
147                .flat_map(|l| {
148                    (l..=n).flat_map(move |r| {
149                        (0..=1).flat_map(move |b| (0..=1).map(move |c| (l, r, b, c)))
150                    })
151                })
152                .collect();
153            let sequences: Vec<_> = if n <= 5 {
154                exhaustive_sequences(updates, 2..=2).collect()
155            } else {
156                (0..10)
157                    .map(|_| {
158                        (0..100)
159                            .map(|_| {
160                                let l = rng.rand(n as u64 + 1) as usize;
161                                let r = l + rng.rand((n - l + 1) as u64) as usize;
162                                (l, r, rng.rand(3) as i32, rng.rand(3) as i32)
163                            })
164                            .collect()
165                    })
166                    .collect()
167            };
168            for uniform in [false, true] {
169                for sequence in &sequences {
170                    let mut values: Vec<_> = if uniform {
171                        vec![M::from(n); n]
172                    } else {
173                        (0..n).map(M::from).collect()
174                    };
175                    let mut seg = if uniform {
176                        DualSegmentTree::<LinearAct<_>>::new(n, M::from(n))
177                    } else {
178                        DualSegmentTree::from_keys(values.iter().copied())
179                    };
180                    for (i, &value) in values.iter().enumerate() {
181                        assert_eq!(seg.get(i), value);
182                    }
183                    for &(l, r, b, c) in sequence {
184                        let (b, c) = (M::from(b), M::from(c));
185                        seg.update(l..r, (b, c));
186                        for value in &mut values[l..r] {
187                            *value = b * *value + c;
188                        }
189                        for (i, &value) in values.iter().enumerate() {
190                            assert_eq!(seg.get(i), value);
191                        }
192                    }
193                    for i in 0..n {
194                        values[i] = M::from(i);
195                        seg.set(i, values[i]);
196                        for (j, &value) in values.iter().enumerate() {
197                            assert_eq!(seg.get(j), value);
198                        }
199                    }
200                }
201            }
202        }
203    }
204}