competitive/data_structure/
segment_tree.rs1use super::{AbelianMonoid, Monoid, RangeBoundsExt};
2use std::{
3 fmt::{self, Debug, Formatter},
4 ops::RangeBounds,
5};
6
7pub struct SegmentTree<M>
8where
9 M: Monoid,
10{
11 n: usize,
12 seg: Vec<M::T>,
13}
14
15impl<M> Clone for SegmentTree<M>
16where
17 M: Monoid,
18{
19 fn clone(&self) -> Self {
20 Self {
21 n: self.n,
22 seg: self.seg.clone(),
23 }
24 }
25}
26
27impl<M> Debug for SegmentTree<M>
28where
29 M: Monoid<T: Debug>,
30{
31 fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
32 f.debug_struct("SegmentTree")
33 .field("n", &self.n)
34 .field("seg", &self.seg)
35 .finish()
36 }
37}
38
39impl<M> SegmentTree<M>
40where
41 M: Monoid,
42{
43 pub fn new(n: usize) -> Self {
44 let seg = vec![M::unit(); 2 * n];
45 Self { n, seg }
46 }
47 pub fn from_vec(v: Vec<M::T>) -> Self {
48 let n = v.len();
49 let mut seg = vec![M::unit(); 2 * n];
50 for (i, x) in v.into_iter().enumerate() {
51 seg[n + i] = x;
52 }
53 for i in (1..n).rev() {
54 seg[i] = M::operate(&seg[2 * i], &seg[2 * i + 1]);
55 }
56 Self { n, seg }
57 }
58 pub fn set(&mut self, k: usize, x: M::T) {
59 assert!(k < self.n);
60 let mut k = k + self.n;
61 self.seg[k] = x;
62 k /= 2;
63 while k > 0 {
64 self.seg[k] = M::operate(&self.seg[2 * k], &self.seg[2 * k + 1]);
65 k /= 2;
66 }
67 }
68 pub fn clear(&mut self, k: usize) {
69 self.set(k, M::unit());
70 }
71 pub fn update(&mut self, k: usize, x: M::T) {
72 assert!(k < self.n);
73 let mut k = k + self.n;
74 self.seg[k] = M::operate(&self.seg[k], &x);
75 k /= 2;
76 while k > 0 {
77 self.seg[k] = M::operate(&self.seg[2 * k], &self.seg[2 * k + 1]);
78 k /= 2;
79 }
80 }
81 pub fn get(&self, k: usize) -> M::T {
82 assert!(k < self.n);
83 self.seg[k + self.n].clone()
84 }
85 pub fn fold<R>(&self, range: R) -> M::T
86 where
87 R: RangeBounds<usize>,
88 {
89 let range = range.to_range_bounded(0, self.n).expect("invalid range");
90 let mut l = range.start + self.n;
91 let mut r = range.end + self.n;
92 let mut vl = M::unit();
93 let mut vr = M::unit();
94 while l < r {
95 if l & 1 != 0 {
96 vl = M::operate(&vl, &self.seg[l]);
97 l += 1;
98 }
99 if r & 1 != 0 {
100 r -= 1;
101 vr = M::operate(&self.seg[r], &vr);
102 }
103 l /= 2;
104 r /= 2;
105 }
106 M::operate(&vl, &vr)
107 }
108 fn partition_point_perfect<P>(
109 &self,
110 mut pos: usize,
111 mut acc: M::T,
112 mut pred: P,
113 ) -> (usize, M::T)
114 where
115 P: FnMut(&M::T) -> bool,
116 {
117 while pos < self.n {
118 pos <<= 1;
119 let nacc = M::operate(&acc, &self.seg[pos]);
120 if pred(&nacc) {
121 acc = nacc;
122 pos += 1;
123 }
124 }
125 (pos - self.n, acc)
126 }
127 fn rpartition_point_perfect<P>(
128 &self,
129 mut pos: usize,
130 mut acc: M::T,
131 mut pred: P,
132 ) -> (usize, M::T)
133 where
134 P: FnMut(&M::T) -> bool,
135 {
136 while pos < self.n {
137 pos = pos * 2 + 1;
138 let nacc = M::operate(&self.seg[pos], &acc);
139 if pred(&nacc) {
140 acc = nacc;
141 pos -= 1;
142 }
143 }
144 (pos - self.n, acc)
145 }
146 pub fn partition_point_acc<P>(&self, left: usize, mut pred: P) -> usize
147 where
148 P: FnMut(&M::T) -> bool,
149 {
150 let mut l = left + self.n;
151 let r = 2 * self.n;
152 let mut k = 0usize;
153 let mut acc = M::unit();
154 while l < r >> k {
155 if l & 1 != 0 {
156 let nacc = M::operate(&acc, &self.seg[l]);
157 if !pred(&nacc) {
158 return self.partition_point_perfect(l, acc, pred).0;
159 }
160 acc = nacc;
161 l += 1;
162 }
163 l >>= 1;
164 k += 1;
165 }
166 for k in (0..k).rev() {
167 let r = r >> k;
168 if r & 1 != 0 {
169 let nacc = M::operate(&acc, &self.seg[r - 1]);
170 if !pred(&nacc) {
171 return self.partition_point_perfect(r - 1, acc, pred).0;
172 }
173 acc = nacc;
174 }
175 }
176 self.n
177 }
178 pub fn rpartition_point_acc<P>(&self, right: usize, mut pred: P) -> usize
179 where
180 P: FnMut(&M::T) -> bool,
181 {
182 let mut l = self.n;
183 let mut r = right + self.n;
184 let mut c = 0usize;
185 let mut k = 0usize;
186 let mut acc = M::unit();
187 while l >> k < r {
188 c <<= 1;
189 if l & (1 << k) != 0 {
190 l += 1 << k;
191 c += 1;
192 }
193 if r & 1 != 0 {
194 r -= 1;
195 let nacc = M::operate(&self.seg[r], &acc);
196 if !pred(&nacc) {
197 return self.rpartition_point_perfect(r, acc, pred).0 + 1;
198 }
199 acc = nacc;
200 }
201 r >>= 1;
202 k += 1;
203 }
204 for k in (0..k).rev() {
205 if c & 1 != 0 {
206 l -= 1 << k;
207 let l = l >> k;
208 let nacc = M::operate(&self.seg[l], &acc);
209 if !pred(&nacc) {
210 return self.rpartition_point_perfect(l, acc, pred).0 + 1;
211 }
212 acc = nacc;
213 }
214 c >>= 1;
215 }
216 0
217 }
218 pub fn as_slice(&self) -> &[M::T] {
219 &self.seg[self.n..]
220 }
221}
222impl<M> SegmentTree<M>
223where
224 M: AbelianMonoid,
225{
226 pub fn fold_all(&self) -> M::T {
227 self.seg[1].clone()
228 }
229}
230
231#[cfg(test)]
232mod tests {
233 use super::*;
234 use crate::{
235 algebra::{AdditiveOperation, MaxOperation},
236 algorithm::SliceBisectExt as _,
237 rand,
238 tools::{NotEmptySegment as Nes, Xorshift},
239 };
240
241 const N: usize = 1_000;
242 const Q: usize = 10_000;
243 const A: i64 = 1_000_000_000;
244
245 #[test]
246 fn test_segment_tree() {
247 let mut rng = Xorshift::default();
248 let mut arr = vec![0; N + 1];
249 let mut seg = SegmentTree::<AdditiveOperation<_>>::new(N);
250 for (k, v) in rng.random_iter((..N, 1..=A)).take(Q) {
251 seg.set(k, v);
252 arr[k + 1] = v;
253 }
254 for i in 0..N {
255 arr[i + 1] += arr[i];
256 }
257 for i in 0..N {
258 for j in i + 1..N + 1 {
259 assert_eq!(seg.fold(i..j), arr[j] - arr[i]);
260 }
261 }
262 for (left, v) in rng.random_iter((..=N, 1..=A * N as i64)).take(Q) {
263 assert_eq!(
264 seg.partition_point_acc(left, |&x| x < v),
265 arr[left + 1..].position_bisect(|&x| x - arr[left] >= v) + left
266 );
267 }
268 for (right, v) in rng.random_iter((..=N, 1..=A)).take(Q) {
269 assert_eq!(
270 seg.rpartition_point_acc(right, |&x| x < v),
271 arr[..right].rposition_bisect(|&x| arr[right] - x >= v)
272 );
273 }
274
275 rand!(rng, mut arr: [-A..=A; N]);
276 let mut seg = SegmentTree::<MaxOperation<_>>::from_vec(arr.clone());
277 for (k, v) in rng.random_iter((..N, -A..=A)).take(Q) {
278 seg.set(k, v);
279 arr[k] = v;
280 }
281 for (l, r) in rng.random_iter(Nes(N)).take(Q) {
282 let res = arr[l..r].iter().max().cloned().unwrap_or_default();
283 assert_eq!(seg.fold(l..r), res);
284 }
285 }
286}