Segment Tree
Segment Tree is a data structure that allows efficient querying and updating of ranges in an array. It is particularly useful for problems involving range queries, such as finding the sum or minimum of elements in a subarray. It is built as a binary tree where each node represents a segment of the array, and the leaves represent individual elements.
Rust Implementation
Generated API reference for the Segment Tree crate. View the rendered API reference →
1//! A generic segment tree.
2//!
3//! `T` is the type of the data in the array; `F` is the merge function
4//! (typically a sum, min, max, or xor closure).
5//!
6//! Both `query` and `update` run in `O(log n)`. The tree is stored in
7//! a flat `Vec<T>` of size `4 * n`, which is the standard bound for an
8//! implicit binary-tree representation.
9
10/// Generic segment tree.
11pub struct SegmentTree<T, F> {
12 size: usize,
13 tree: Vec<T>,
14 default: T,
15 merge: F,
16}
17
18impl<T, F> SegmentTree<T, F>
19where
20 T: Clone + Copy,
21 F: Fn(T, T) -> T,
22{
23 /// Build a segment tree over `arr` using `default` as the identity
24 /// for `merge` and `merge` itself as the associative combining op.
25 pub fn new(arr: &[T], default: T, merge: F) -> Self {
26 let size = arr.len();
27 #[allow(unused_mut)]
28 let mut tree = vec![default; 4 * size];
29 let mut seg_tree = SegmentTree {
30 size,
31 tree,
32 default,
33 merge,
34 };
35 seg_tree.build(arr, 0, 0, size - 1);
36 seg_tree
37 }
38
39 fn build(&mut self, arr: &[T], node: usize, start: usize, end: usize) {
40 if start == end {
41 self.tree[node] = arr[start];
42 return;
43 }
44 let mid = (start + end) / 2;
45 self.build(arr, 2 * node + 1, start, mid);
46 self.build(arr, 2 * node + 2, mid + 1, end);
47 self.tree[node] = (self.merge)(self.tree[2 * node + 1], self.tree[2 * node + 2]);
48 }
49
50 /// Query `merge` over the closed range `[l, r]`.
51 pub fn query(&self, l: usize, r: usize) -> T {
52 self.query_recursive(0, 0, self.size - 1, l, r)
53 }
54
55 fn query_recursive(&self, node: usize, start: usize, end: usize, l: usize, r: usize) -> T {
56 if r < start || end < l {
57 return self.default;
58 }
59 if l <= start && end <= r {
60 return self.tree[node];
61 }
62 let mid = (start + end) / 2;
63 let p1 = self.query_recursive(2 * node + 1, start, mid, l, r);
64 let p2 = self.query_recursive(2 * node + 2, mid + 1, end, l, r);
65 (self.merge)(p1, p2)
66 }
67
68 /// Point update: set `arr[idx] = val` and recompute affected nodes.
69 pub fn update(&mut self, idx: usize, val: T) {
70 self.update_recursive(0, 0, self.size - 1, idx, val);
71 }
72
73 fn update_recursive(&mut self, node: usize, start: usize, end: usize, idx: usize, val: T) {
74 if start == end {
75 self.tree[node] = val;
76 return;
77 }
78 let mid = (start + end) / 2;
79 if start <= idx && idx <= mid {
80 self.update_recursive(2 * node + 1, start, mid, idx, val);
81 } else {
82 self.update_recursive(2 * node + 2, mid + 1, end, idx, val);
83 }
84 self.tree[node] = (self.merge)(self.tree[2 * node + 1], self.tree[2 * node + 2]);
85 }
86}