引入:线段树是一种二叉树形数据结构,用来高效处理数组的区间查询与区间修改。它可以在O(logn)的时间内完成区间求和、区间最值等操作,功能比树状数组更强大,但同时代码也更复杂一些。
原理:线段树的基本思想就是将数组a[1…n]不断二分,构建一棵二叉树:
叶子节点:代表单个元素 a[i]
内部节点:代表一段区间 [l,r],它的值由左右儿子合并得到
这样,任意一段查询区间[L,R]都可以用O(logn)个节点的值合并出来;更新一个位置只需从叶子走到根,更新O(logn)个节点。
模板主体:
template<class T>
class SegmentTree
{
private:
//用于存储区间节点的数组
vector<T> m_tree;
//懒惰数组
vector<T> m_lazy;
//原数组长度
int m_size;
//初始化函数,将原数组转化为线段树的二叉树结构的数组
void build(const vector<T>& arr, int node, int start, int end);
//处理懒惰标记
void pushDown(int node, int start, int end);
//区间修改
void updateRange(int node, int start, int end, int l, int r, T val);
//区间查询
T queryRange(int node, int start, int end, int l, int r);
public:
SegmentTree(const vector<T>& arr);
void Update(int l, int r, T val);
T Query(int l, int r);
};1、初始化函数build
作用:初始化线段树,使每个节点保存对应区间的和。时间复杂度 O(n)。
template<class T>
void SegmentTree<T>::build(const vector<T>& arr, int node, int start, int end) {
if (start == end) {
m_tree[node] = arr[start];
return;
}
int mid = (start + end) >> 1;
int leftChild = node * 2;
int rightChild = node * 2 + 1;
build(arr, leftChild, start, mid);
build(arr, rightChild, mid + 1, end);
m_tree[node] = m_tree[leftChild] + m_tree[rightChild];
}2、处理懒惰标记的函数pushDown
原理:当需要访问某节点的儿子时,如果该节点有尚未处理的区间加法(m_lazy[node] != 0),则必须把懒惰值下传给两个儿子,否则儿子保存的和是错误的。
作用:保证在查询或更新进入子节点前,子节点的值是最新的,使得区间修改的复杂度可以控制在 O(logn)。
template<class T>
void SegmentTree<T>::pushDown(int node, int start, int end) {
if (m_lazy[node] != 0) {
int mid = (start + end) >> 1;
int leftChild = node * 2;
m_lazy[leftChild] += m_lazy[node];
m_tree[leftChild] += (mid - start + 1) * m_lazy[node];
int rightChild = node * 2 + 1;
m_lazy[rightChild] += m_lazy[node];
m_tree[rightChild] += (end - mid) * m_lazy[node];
//将标记传给儿子时,自身标记变为0
m_lazy[node] = 0;
}
}3、区间修改函数updateRange
作用:将区间 [l, r] 每个元素加上 val,利用懒惰标记避免每次更新都深入到叶子,实现 O(logn)时间。
template<class T>
void SegmentTree<T>::updateRange(int node, int start, int end, int l, int r, T val) {
if (start > r || end < l) {
return;
}
if (start >= l && end <= r) {
m_tree[node] += val * (end - start + 1);
m_lazy[node] += val;
return;
}
pushDown(node, start, end);
int mid = (start + end) >> 1;
int leftChild = node * 2;
int rightChild = node * 2 + 1;
updateRange(leftChild, start, mid, l, r, val);
updateRange(rightChild, mid + 1, end, l, r, val);
m_tree[node] = m_tree[leftChild] + m_tree[rightChild];
}4、区间查询函数queryRange
作用:返回区间 [l, r] 的元素和,时间复杂度 O(logn)。
template<class T>
T SegmentTree<T>::queryRange(int node, int start, int end, int l, int r) {
if (start > r || end < l) {
return 0;
}
if (start >= l && end <= r) {
return m_tree[node];
}
pushDown(node, start, end);
int mid = (start + end) >> 1;
int leftChild = node * 2;
int rightChild = node * 2 + 1;
T leftSum = queryRange(leftChild, start, mid, l, r);
T rightSum = queryRange(rightChild, mid + 1, end, l, r);
return leftSum + rightSum;
}5、构造函数以及对外开放调用函数
template<class T>
SegmentTree<T>::SegmentTree(const vector<T>& arr) {
//由于数组下标从1开始,所以元素个数为容器大小-1
m_size = arr.size() - 1;
//设置容量为原数组的大小的四倍,即可装下所有的区间节点
m_tree.resize(m_size * 4);
m_lazy.resize(m_size * 4, 0);
build(arr, 1, 1, m_size);
}
template<class T>
void SegmentTree<T>::Update(int l, int r, T val) {
updateRange(1, 1, m_size, l, r, val);
}
template<class T>
T SegmentTree<T>::Query(int l, int r) {
return queryRange(1, 1, m_size, l, r);
}完整模板(方便复制)
template<class T>
class SegmentTree
{
private:
//数组下标从1开始!
vector<T> m_tree;
vector<T> m_lazy;
int m_size;
void build(const vector<T>& arr, int node, int start, int end);
void pushDown(int node, int start, int end);
void updateRange(int node, int start, int end, int l, int r, T val);
T queryRange(int node, int start, int end, int l, int r);
public:
SegmentTree(const vector<T>& arr);
void Update(int l, int r, T val);
T Query(int l, int r);
};
template<class T>
void SegmentTree<T>::build(const vector<T>& arr, int node, int start, int end) {
if (start == end) {
m_tree[node] = arr[start];
return;
}
int mid = (start + end) >> 1;
int leftChild = node * 2;
int rightChild = node * 2 + 1;
build(arr, leftChild, start, mid);
build(arr, rightChild, mid + 1, end);
m_tree[node] = m_tree[leftChild] + m_tree[rightChild];
}
template<class T>
void SegmentTree<T>::pushDown(int node, int start, int end) {
if (m_lazy[node] != 0) {
int mid = (start + end) >> 1;
int leftChild = node * 2;
m_lazy[leftChild] += m_lazy[node];
m_tree[leftChild] += (mid - start + 1) * m_lazy[node];
int rightChild = node * 2 + 1;
m_lazy[rightChild] += m_lazy[node];
m_tree[rightChild] += (end - mid) * m_lazy[node];
m_lazy[node] = 0;
}
}
template<class T>
void SegmentTree<T>::updateRange(int node, int start, int end, int l, int r, T val) {
if (start > r || end < l) {
return;
}
if (start >= l && end <= r) {
m_tree[node] += val * (end - start + 1);
m_lazy[node] += val;
return;
}
pushDown(node, start, end);
int mid = (start + end) >> 1;
int leftChild = node * 2;
int rightChild = node * 2 + 1;
updateRange(leftChild, start, mid, l, r, val);
updateRange(rightChild, mid + 1, end, l, r, val);
m_tree[node] = m_tree[leftChild] + m_tree[rightChild];
}
template<class T>
T SegmentTree<T>::queryRange(int node, int start, int end, int l, int r) {
if (start > r || end < l) {
return 0;
}
if (start >= l && end <= r) {
return m_tree[node];
}
pushDown(node, start, end);
int mid = (start + end) >> 1;
int leftChild = node * 2;
int rightChild = node * 2 + 1;
T leftSum = queryRange(leftChild, start, mid, l, r);
T rightSum = queryRange(rightChild, mid + 1, end, l, r);
return leftSum + rightSum;
}
template<class T>
SegmentTree<T>::SegmentTree(const vector<T>& arr) {
m_size = arr.size() - 1;
m_tree.resize(m_size * 4);
m_lazy.resize(m_size * 4, 0);
build(arr, 1, 1, m_size);
}
template<class T>
void SegmentTree<T>::Update(int l, int r, T val) {
updateRange(1, 1, m_size, l, r, val);
}
template<class T>
T SegmentTree<T>::Query(int l, int r) {
return queryRange(1, 1, m_size, l, r);
}
评论