competitive/data_structure/
dual_segment_tree.rs1use 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}