competitive/data_structure/
lazy_segment_tree.rs1use super::{LazyMapMonoid, RangeBoundsExt};
2use std::{
3 fmt::{self, Debug, Formatter},
4 mem::replace,
5 ops::RangeBounds,
6};
7
8pub struct LazySegmentTree<M>
9where
10 M: LazyMapMonoid,
11{
12 len: usize,
13 n: usize,
14 seg: Vec<M::Agg>,
15 lazy: Vec<M::Act>,
16}
17
18impl<M> Clone for LazySegmentTree<M>
19where
20 M: LazyMapMonoid,
21{
22 fn clone(&self) -> Self {
23 Self {
24 len: self.len,
25 n: self.n,
26 seg: self.seg.clone(),
27 lazy: self.lazy.clone(),
28 }
29 }
30}
31
32impl<M> Debug for LazySegmentTree<M>
33where
34 M: LazyMapMonoid<Agg: Debug, Act: Debug>,
35{
36 fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
37 f.debug_struct("LazySegmentTree")
38 .field("len", &self.len)
39 .field("n", &self.n)
40 .field("seg", &self.seg)
41 .field("lazy", &self.lazy)
42 .finish()
43 }
44}
45
46impl<M> LazySegmentTree<M>
47where
48 M: LazyMapMonoid,
49{
50 pub fn new(len: usize) -> Self {
51 let n = len.next_power_of_two();
52 let seg = vec![M::agg_unit(); 2 * n];
53 let lazy = vec![M::act_unit(); n];
54 Self { len, n, seg, lazy }
55 }
56 pub fn from_vec(v: Vec<M::Agg>) -> Self {
57 let len = v.len();
58 let n = len.next_power_of_two();
59 let mut seg = vec![M::agg_unit(); 2 * n];
60 for (i, x) in v.into_iter().enumerate() {
61 seg[i + n] = x;
62 }
63 for i in (1..n).rev() {
64 seg[i] = M::agg_operate(&seg[2 * i], &seg[2 * i + 1]);
65 }
66 let lazy = vec![M::act_unit(); n];
67 Self { len, n, seg, lazy }
68 }
69 pub fn from_keys(keys: impl ExactSizeIterator<Item = M::Key>) -> Self {
70 let len = keys.len();
71 let n = len.next_power_of_two();
72 let mut seg = vec![M::agg_unit(); 2 * n];
73 for (i, key) in keys.enumerate() {
74 seg[i + n] = M::single_agg(&key);
75 }
76 for i in (1..n).rev() {
77 seg[i] = M::agg_operate(&seg[2 * i], &seg[2 * i + 1]);
78 }
79 let lazy = vec![M::act_unit(); n];
80 Self { len, n, seg, lazy }
81 }
82 #[inline]
83 fn update_at(&mut self, k: usize, x: &M::Act) {
84 if M::is_act_unit(x) {
85 return;
86 }
87 let nx = M::act_agg(&self.seg[k], x);
88 if k < self.n {
89 self.lazy[k] = M::act_operate(&self.lazy[k], x);
90 }
91 if let Some(nx) = nx {
92 self.seg[k] = nx;
93 } else if k < self.n {
94 self.propagate_at(k);
95 self.recalc_at(k);
96 } else {
97 panic!("act failed on leaf");
98 }
99 }
100 #[inline]
101 fn recalc_at(&mut self, k: usize) {
102 self.seg[k] = M::agg_operate(&self.seg[2 * k], &self.seg[2 * k + 1]);
103 }
104 #[inline]
105 fn propagate_at(&mut self, k: usize) {
106 debug_assert!(k < self.n);
107 let x = replace(&mut self.lazy[k], M::act_unit());
108 if M::is_act_unit(&x) {
109 return;
110 }
111 self.update_at(2 * k, &x);
112 self.update_at(2 * k + 1, &x);
113 }
114 #[inline]
115 fn propagate(&mut self, k: usize) {
116 for i in (1..=self.n.trailing_zeros()).rev() {
117 self.propagate_at(k >> i);
118 }
119 }
120 #[inline]
121 fn recalc(&mut self, mut k: usize) {
122 while k > 1 {
123 k >>= 1;
124 self.recalc_at(k);
125 }
126 }
127 pub fn update<R>(&mut self, range: R, x: M::Act)
128 where
129 R: RangeBounds<usize>,
130 {
131 let range = range.to_range_bounded(0, self.len).expect("invalid range");
132 if range.is_empty() || M::is_act_unit(&x) {
133 return;
134 }
135 let mut a = range.start + self.n;
136 let mut b = range.end + self.n;
137 for i in (1..=self.n.trailing_zeros()).rev() {
138 if (a >> i) << i != a {
139 self.propagate_at(a >> i);
140 }
141 if (b >> i) << i != b {
142 self.propagate_at((b - 1) >> i);
143 }
144 }
145 while a < b {
146 if a & 1 != 0 {
147 self.update_at(a, &x);
148 a += 1;
149 }
150 if b & 1 != 0 {
151 b -= 1;
152 self.update_at(b, &x);
153 }
154 a /= 2;
155 b /= 2;
156 }
157 let a = range.start + self.n;
158 let b = range.end + self.n;
159 for i in 1..=self.n.trailing_zeros() {
160 if (a >> i) << i != a {
161 self.recalc_at(a >> i);
162 }
163 if (b >> i) << i != b {
164 self.recalc_at((b - 1) >> i);
165 }
166 }
167 }
168 pub fn fold<R>(&mut self, range: R) -> M::Agg
169 where
170 R: RangeBounds<usize>,
171 {
172 let range = range.to_range_bounded(0, self.len).expect("invalid range");
173 if range.is_empty() {
174 return M::agg_unit();
175 }
176 if let Some(result) = (|| {
177 let mut left_index = range.start + self.n - 1;
178 let mut right_index = range.end + self.n;
179 let mut left = M::agg_unit();
180 let mut right = M::agg_unit();
181 let mut has_left = false;
182 let mut has_right = false;
183 for _ in 0..(left_index ^ right_index).ilog2() {
184 if left_index & 1 == 0 {
185 left = M::agg_operate(&left, &self.seg[left_index ^ 1]);
186 has_left = true;
187 }
188 if right_index & 1 != 0 {
189 right = M::agg_operate(&self.seg[right_index ^ 1], &right);
190 has_right = true;
191 }
192 left_index >>= 1;
193 right_index >>= 1;
194 if has_left {
195 left = M::act_agg(&left, &self.lazy[left_index])?;
196 }
197 if has_right && right_index < self.n {
198 right = M::act_agg(&right, &self.lazy[right_index])?;
199 }
200 }
201 let mut result = M::agg_operate(&left, &right);
202 while left_index > 1 {
203 left_index >>= 1;
204 result = M::act_agg(&result, &self.lazy[left_index])?;
205 }
206 Some(result)
207 })() {
208 return result;
209 }
210 let mut l = range.start + self.n;
211 let mut r = range.end + self.n;
212 self.propagate(l);
213 self.propagate(r - 1);
214 let mut vl = M::agg_unit();
215 let mut vr = M::agg_unit();
216 while l < r {
217 if l & 1 != 0 {
218 vl = M::agg_operate(&vl, &self.seg[l]);
219 l += 1;
220 }
221 if r & 1 != 0 {
222 r -= 1;
223 vr = M::agg_operate(&self.seg[r], &vr);
224 }
225 l /= 2;
226 r /= 2;
227 }
228 M::agg_operate(&vl, &vr)
229 }
230 pub fn set(&mut self, k: usize, x: M::Agg) {
231 assert!(k < self.len);
232 let k = k + self.n;
233 self.propagate(k);
234 self.seg[k] = x;
235 self.recalc(k);
236 }
237 pub fn get(&mut self, k: usize) -> M::Agg {
238 self.fold(k..k + 1)
239 }
240 pub fn fold_all(&self) -> M::Agg {
241 self.seg[1].clone()
242 }
243 pub fn partition_point_acc<P>(&mut self, left: usize, mut pred: P) -> usize
244 where
245 P: FnMut(&M::Agg) -> bool,
246 {
247 let mut acc = M::agg_unit();
248 if left == self.len {
249 return self.len;
250 }
251 let mut k = left + self.n;
252 self.propagate(k);
253 loop {
254 while k & 1 == 0 {
255 k >>= 1;
256 }
257 let nacc = M::agg_operate(&acc, &self.seg[k]);
258 if !pred(&nacc) {
259 while k < self.n {
260 self.propagate_at(k);
261 k <<= 1;
262 let nacc = M::agg_operate(&acc, &self.seg[k]);
263 if pred(&nacc) {
264 acc = nacc;
265 k += 1;
266 }
267 }
268 return k - self.n;
269 }
270 acc = nacc;
271 k += 1;
272 if k.is_power_of_two() {
273 return self.len;
274 }
275 }
276 }
277 pub fn rpartition_point_acc<P>(&mut self, right: usize, mut pred: P) -> usize
278 where
279 P: FnMut(&M::Agg) -> bool,
280 {
281 let mut acc = M::agg_unit();
282 if right == 0 {
283 return 0;
284 }
285 let mut k = right + self.n;
286 self.propagate(k - 1);
287 loop {
288 k -= 1;
289 while k > 1 && k & 1 != 0 {
290 k >>= 1;
291 }
292 let nacc = M::agg_operate(&self.seg[k], &acc);
293 if !pred(&nacc) {
294 while k < self.n {
295 self.propagate_at(k);
296 k = 2 * k + 1;
297 let nacc = M::agg_operate(&self.seg[k], &acc);
298 if pred(&nacc) {
299 acc = nacc;
300 k -= 1;
301 }
302 }
303 return k + 1 - self.n;
304 }
305 acc = nacc;
306 if k.is_power_of_two() {
307 return 0;
308 }
309 }
310 }
311}
312
313#[cfg(test)]
314mod tests {
315 use super::*;
316 use crate::{
317 algebra::{
318 RangeChminChmaxAdd, RangeMaxRangeUpdate, RangeSumRangeAdd, RangeSumRangeChminChmaxAdd,
319 },
320 num::Saturating,
321 rand,
322 tools::{NotEmptySegment, Xorshift},
323 };
324
325 const N: usize = 1_000;
326 const Q: usize = 20_000;
327 const A: i64 = 1_000_000_000;
328
329 #[test]
330 fn test_lazy_segment_tree() {
331 let mut rng = Xorshift::default();
332 rand!(rng, mut arr: [-A..A; N]);
334 let mut seg =
335 LazySegmentTree::<RangeSumRangeAdd<_>>::from_vec(arr.iter().map(|&a| (a, 1)).collect());
336 for _ in 0..Q {
337 rand!(rng, (l, r): NotEmptySegment(N));
338 match rng.rand(3) {
339 0 => {
340 rand!(rng, x: -A..A);
342 seg.update(l..r, x);
343 for a in arr[l..r].iter_mut() {
344 *a += x;
345 }
346 }
347 1 => {
348 rand!(rng, k: 0..N, x: -A..A);
350 seg.set(k, (x, 1));
351 arr[k] = x;
352 }
353 _ => {
354 let res = arr[l..r].iter().sum();
356 assert_eq!(seg.fold(l..r).0, res);
357 }
358 }
359 rand!(rng, k: 0..N);
360 assert_eq!(seg.get(k).0, arr[k]);
361 assert_eq!(seg.fold_all().0, arr.iter().sum());
362 }
363
364 rand!(rng, mut arr: [-A..A; N]);
366 let mut seg = LazySegmentTree::<RangeMaxRangeUpdate<_>>::from_vec(arr.clone());
367 for _ in 0..Q {
368 rand!(rng, ty: 0..5, (l, r): NotEmptySegment(N));
369 match ty {
370 0 => {
371 rand!(rng, x: -A..A);
373 seg.update(l..r, Some(x));
374 arr[l..r].iter_mut().for_each(|a| *a = x);
375 }
376 1 => {
377 let res = arr[l..r].iter().max().cloned().unwrap_or_default();
379 assert_eq!(seg.fold(l..r), res);
380 }
381 2 => {
382 rand!(rng, left: ..=N, x: -A..A);
384 assert_eq!(
385 seg.partition_point_acc(left, |&d| d < x),
386 arr[left..]
387 .iter()
388 .scan(i64::MIN, |acc, &a| {
389 *acc = a.max(*acc);
390 Some(*acc)
391 })
392 .position(|acc| acc >= x)
393 .map_or(N, |i| i + left),
394 );
395 }
396 3 => {
397 rand!(rng, right: ..=N, x: -A..A);
399 assert_eq!(
400 seg.rpartition_point_acc(right, |&d| d < x),
401 arr[..right]
402 .iter()
403 .rev()
404 .scan(i64::MIN, |acc, &a| {
405 *acc = a.max(*acc);
406 Some(*acc)
407 })
408 .position(|acc| acc >= x)
409 .map_or(0, |i| right - i),
410 );
411 }
412 _ => {
413 rand!(rng, k: 0..N, x: -A..A);
415 seg.set(k, x);
416 arr[k] = x;
417 }
418 }
419 rand!(rng, k: 0..N);
420 assert_eq!(seg.get(k), arr[k]);
421 assert_eq!(seg.fold_all(), *arr.iter().max().unwrap());
422 }
423
424 let mut arr = rng
426 .random_iter(-1_000..=1_000)
427 .map(Saturating)
428 .take(N)
429 .collect::<Vec<_>>();
430 let mut seg =
431 LazySegmentTree::<RangeSumRangeChminChmaxAdd<_>>::from_keys(arr.iter().copied());
432 for _ in 0..Q {
433 rand!(rng, ty: 0..4, (l, r): NotEmptySegment(N), x: -1_000..=1_000);
434 let x = Saturating(x);
435 match ty {
436 0 => {
437 seg.update(l..r, RangeChminChmaxAdd::chmin(x));
438 arr[l..r].iter_mut().for_each(|a| *a = (*a).min(x));
439 }
440 1 => {
441 seg.update(l..r, RangeChminChmaxAdd::chmax(x));
442 arr[l..r].iter_mut().for_each(|a| *a = (*a).max(x));
443 }
444 2 => {
445 seg.update(l..r, RangeChminChmaxAdd::add(x));
446 arr[l..r].iter_mut().for_each(|a| *a += x);
447 }
448 _ => assert_eq!(
449 seg.fold(l..r).sum,
450 arr[l..r].iter().copied().sum::<Saturating<i64>>()
451 ),
452 }
453 }
454 }
455}