引入:线段树是一种二叉树形数据结构,用来高效处理数组的区间查询与区间修改。它可以在O(log⁡n)的时间内完成区间求和、区间最值等操作,功能比树状数组更强大,但同时代码也更复杂一些。

原理:线段树的基本思想就是将数组a[1…n]不断二分,构建一棵二叉树:

  • 叶子节点:代表单个元素 a[i]

  • 内部节点:代表一段区间 [l,r],它的值由左右儿子合并得到

这样,任意一段查询区间[L,R]都可以用O(log⁡n)个节点的值合并出来;更新一个位置只需从叶子走到根,更新O(log⁡n)个节点。

模板题:P3372 【模板】线段树 1 - 洛谷

模板主体:

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(log⁡n)

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(log⁡n)时间。

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(log⁡n)

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);
}