平衡树

· · 算法·理论

二叉搜索树(BST)

二叉搜索树的性质

二叉搜索树具有以下性质:

  1. 每个节点对应一个权值。

  2. 若当前节点左子树非空,则左子节点的权值小于当前节点的权值。

  3. 若当前节点右子树非空,则右子节点的权值大于当前节点的权值。

  4. 左右子树均为二叉搜索树。

  5. 对二叉搜索树进行中序遍历,能得到一个从小到大排列的序列。

    二叉搜索树的遍历

    要在二叉搜索树中输出一个从小到大排列的序列,只需要输出二叉搜索树的中序遍历。对于每个节点,先输出左子树,再输出当前节点的权值,最后输出右子树即可。

    二叉搜索树的查找

    在二叉搜索树中查找是否有特定的值,方法如下:

  6. 从根节点开始查找。

  7. 如果当前节点为要查找的值,则存在这个值。

  8. 如果当前节点小于要查找的值,则在右子树中查找。

  9. 如果当前节点大于要查找的值,则在左子树中查找。

  10. 如果当前节点为空节点,则不存在这个值。

    在二叉搜索树中查找最值

    以最小值为例,方法如下:

  11. 从根节点开始查找。

  12. 如果当前节点存在左子树,则在左子树中查找。

  13. 如果当前节点不存在左子树,则当前节点为最小值。

查找最大值类似。

在二叉搜索树中插入节点

和查找特定值差不多,按查找当前值的方法进入那个空节点,然后将空节点修改为要插入的节点,左右子树为空。

在二叉搜索树中删除节点

删除特定节点,需要分情况讨论:

  1. 该节点是叶节点,直接删除即可。

  2. 该节点只有一个子节点,将子节点替换即可。

  3. 该节点有两个子节点,一般用右子树的最小值替换当前节点然后把右子树的最小值删除。

    二叉平衡树

    二叉平衡树的性质

    二叉平衡树是一种特殊的二叉搜索树,满足以下性质:

  4. 左右子树均为二叉平衡树。

  5. 左右子树的高度差绝对值小于等于 1

    二叉平衡树的旋转

    在插入/删除节点后,通常需要通过旋转维持树的平衡。

    左旋

    将右子节点变成当前节点的父节点,将右子节点的左子节点变成当前节点的右子节点。

    右旋

    将左子节点变成当前节点的父节点,将左子节点的右子节点变成当前节点的左子节点。

    Treap

    Treap 的性质

    Treap 是由二叉搜索树与堆数据结构组合而成,每个节点有 valdata 信息,分别表示每个节点对应的权值及优先级,data 一般为随机值。其中 val 满足二叉搜索树的性质,data 满足堆性质。

    Treap 的插入

    和在二叉搜索树中插入类似,只是新的 data 会改变堆的性质,所以在插入完后需要通过旋转调整为堆。具体操作如下:

  6. 当左子节点的 data 小于当前节点的 data 时,将以当前节点为根的子树进行右旋操作。

  7. 当右子节点的 data 小于当前节点的 data 时,将以当前节点为根的子树进行左旋操作。

    Treap 的删除

    当要删除的节点孩子数量小于等于 1 时,和二叉搜索树的删除一样。当有 2 个孩子时,需要通过旋转使得要删除的节点只有 01 个孩子,具体操作如下:

  8. 当左子节点的 data 小于右子节点的 data 时,进行左旋操作,然后删除右子节点

  9. 当右子节点的 data 小于左子节点的 data 时,进行右旋操作,然后删除左子节点。

    Treap 查询排名第 k 的元素

    不难发现,记录 sz_x 为以 x 号节点为根的子树大小,cnt_xx 号节点出现的次数,ls_xx 号节点的左子节点,则 x 号节点在子树中的排名为 [sz_{ls_x}+1,sz_{ls_x}+cnt_x]。所以查询排名时,若 sz_{ls_x}+1\le k \le sz_{ls_x}+cnt_x,则答案为 x。若 k\le sz_{ls_x},则在左子树中找排名为 k 的元素。若 k>sz_{ls_x}+cnt_x,则在右子树中找排名为 k-sz_{ls_x}-cnt_x 的值。

    Treap 查询小于 k 的个数

  10. 若走到了空节点,则答案为 0

  11. 若当前节点的 valk,则答案为 sz_{ls_x}

  12. 若当前节点的 val 大于 k,则答案为左子树中小于 k 的个数。

  13. 若当前节点的 val 小于 k,则答案为右子树中小于 k 的个数 +sz_{ls_x}+cnt_x

    Treap 中求前驱后继节点

    先求出当前节点的排名,即为小于当前数的个数加一,记为 rk。则前驱为排名为 rk-1 的数,后继为排名为 rk+1 的数。

    FHQ-Treap

    FHQ-Treap 是基于合并与分裂实现的。

    FHQ-Treap 的分裂

    分裂也就是将一个 Treap 变成两个 Treap,记作 lr,其中 l 中所有元素 <kr 中所有元素 \ge k。每走到一个节点,需要分以下两种情况讨论:

  14. 当前节点的 val 小于 k,则将左子树都划进 l,分裂右子树,将分裂后右子树 l 的根节点设为当前节点的右子节点。

  15. 当前节点的 val 大于 k,则将右子树都划进 r,分裂左子树,将分裂后左子树 r 的根节点设为当前节点的左子节点。

    FHQ-Treap 的合并

    对于 lr 两棵子树进行合并,一定要满足 l 中所有元素均小于 r 中所有元素。为了防止被卡,我们使用一个随机权值,由随机权值大的合并到随机权值小的。具体过程如下:

  16. 有一个子树为空,合并结果为另外一个。

  17. ### FHQ-Treap 的插入 这非常简单,按要插入的值分裂成 $l$ 和 $r$,然后合并 $l$、插入的值、$r$,就可以了。 ### FHQ-Treap 的删除 先按删除的值分裂成 $a$ 和 $b$,再将 $b$ 按删除的值加一分为 $c$ 和 $d$,再合并 $a$ 和 $d$ 就可以了。 ### FHQ-Treap 查询排名 直接将树按查询排名的值分裂,然后输出 $l$ 的大小加一即可。 ### FHQ-Treap 查询排名第 $k$ 的数 和普通 Treap 一样,此处不赘述。 ### FHQ-Treap 查询前驱后继 仍然和普通 Treap 一样。 ### FHQ-Treap 解决区间问题 我们可以借助线段树的思想,将需要进行区间操作的区间从原树中分裂出来,打上懒标记,再合并进去就行了。但是在分裂和合并前,需要将懒标记下传到子节点,才能解决。 ## 例题 ### [P3369 【模板】普通平衡树](https://www.luogu.com.cn/problem/P3369) 模板题,上面已经讲过了。

普通 Treap 做法:

#include <bits/stdc++.h>
using namespace std;
struct treap{
    int val, ls, rs, sz, data, cnt;
} tr[16000005];
int n, opt, x, tot = 0, root = 0;
mt19937 rd(time(0));
void zig(int &k){
    int y = tr[k].ls;
    tr[k].ls = tr[y].rs;
    tr[y].rs = k;
    tr[y].sz = tr[k].sz;
    tr[k].sz = tr[tr[k].ls].sz + tr[tr[k].rs].sz + tr[k].cnt;
    k = y;
}
void zag(int &k){
    int y = tr[k].rs;
    tr[k].rs = tr[y].ls;
    tr[y].ls = k;
    tr[y].sz = tr[k].sz;
    tr[k].sz = tr[tr[k].ls].sz + tr[tr[k].rs].sz + tr[k].cnt;
    k = y;
}
void ins(int &x, int k){
    if(x == 0){
        x = ++tot;
        tr[x].val = k;
        tr[x].sz = 1;
        tr[x].cnt = 1;
        tr[x].data = rd();
        tr[x].ls = tr[x].rs = 0;
        return;
    }
    if(tr[x].val == k){
        tr[x].cnt++;
        tr[x].sz++;
        return;
    }
    if(k < tr[x].val) ins(tr[x].ls, k);
    else ins(tr[x].rs, k);
    if(tr[x].ls && tr[tr[x].ls].data < tr[x].data) zig(x);
    if(tr[x].rs && tr[tr[x].rs].data < tr[x].data) zag(x);
    tr[x].sz = tr[tr[x].ls].sz + tr[tr[x].rs].sz + tr[x].cnt;
}
void del(int &x, int k){
    if(x == 0) return;
    if(tr[x].val == k){
        if(tr[x].cnt > 1){
            tr[x].cnt--;
            tr[x].sz--;
            return;
        }
        if(!tr[x].ls || !tr[x].rs){
            x = tr[x].ls | tr[x].rs;
        }
        else{
            if(tr[tr[x].ls].data < tr[tr[x].rs].data){
                zig(x);
                del(x, k);
            }
            else{
                zag(x);
                del(x, k);
            }
        }
    }
    else if(k < tr[x].val) del(tr[x].ls, k);
    else del(tr[x].rs, k);
    if(x) tr[x].sz = tr[tr[x].ls].sz + tr[tr[x].rs].sz + tr[x].cnt;
}
int find_by_order(int p, int k){
    if(tr[tr[p].ls].sz + 1 <= k && k <= tr[tr[p].ls].sz + tr[p].cnt) return tr[p].val;
    if(k <= tr[tr[p].ls].sz) return find_by_order(tr[p].ls, k);
    return find_by_order(tr[p].rs, k - tr[tr[p].ls].sz - tr[p].cnt);
}
int order_of_key(int p, int k){
    if(p == 0) return 1;
    if(k == tr[p].val) return tr[tr[p].ls].sz + 1;
    if(k < tr[p].val) return order_of_key(tr[p].ls, k);
    return tr[tr[p].ls].sz + tr[p].cnt + order_of_key(tr[p].rs, k);
}
int main(){
    cin >> n;
    while (n--){
        cin >> opt >> x;
        if(opt == 1) ins(root, x);
        if(opt == 2) del(root, x);
        if(opt == 3) cout << order_of_key(root, x) << '\n';
        if(opt == 4) cout << find_by_order(root, x) << '\n';
        if(opt == 5){
            int rk = order_of_key(root, x) - 1;
            cout << find_by_order(root, rk) << '\n';
        }
        if(opt == 6){
            int rk = order_of_key(root, x + 1);
            cout << find_by_order(root, rk) << '\n';
        }
    }
    return 0;
}

P3391 【模版】文艺平衡树

这是一道典型的 FHQ-Treap 解决区间问题的例子。这是区间翻转,所以懒标记是取反的,因为翻转两次和没翻转是一样的。

代码如下:

#include <bits/stdc++.h>
using namespace std;
struct Treap{
    int val, ls, rs, data, sz, tag;
}tr[100005];
mt19937 rd(time(0));
int n, m, tot, root;
void push_down(int x){
    if(!tr[x].tag) return;
    swap(tr[x].ls, tr[x].rs);
    if(tr[x].ls) tr[tr[x].ls].tag ^= 1;
    if(tr[x].rs) tr[tr[x].rs].tag ^= 1;
    tr[x].tag = 0;
}
pair<int, int> split(int x, int k){
    if(!x) return {0, 0};
    push_down(x);
    if(tr[tr[x].ls].sz < k){
        auto [u, v] = split(tr[x].rs, k - tr[tr[x].ls].sz - 1);
        tr[x].rs = u;
        tr[x].sz = tr[tr[x].ls].sz + 1 + tr[tr[x].rs].sz;
        return {x, v};
    }
    auto [u, v] = split(tr[x].ls, k);
    tr[x].ls = v;
    tr[x].sz = tr[tr[x].ls].sz + 1 + tr[tr[x].rs].sz;
    return {u, x};
}
int merge(int l, int r){
    if(!l || !r) return l | r;
    push_down(l);
    push_down(r);
    if(tr[l].data < tr[r].data){
        tr[l].rs = merge(tr[l].rs, r);
        tr[l].sz = tr[tr[l].ls].sz + 1 + tr[tr[l].rs].sz;
        return l;
    }
    tr[r].ls = merge(l, tr[r].ls);
    tr[r].sz = tr[tr[r].ls].sz + 1 + tr[tr[r].rs].sz;
    return r;
}
int query(int x){
    push_down(x);
    auto [l, midr] = split(root, x - 1);
    auto [mid, r] = split(midr, 1);
    int ans = tr[mid].val;
    root = merge(merge(l, mid), r);
    return ans;
}
int main(){
    cin >> n >> m;
    for(int i = 1; i <= n; i++){
        tr[++tot] = {i, 0, 0, rd(), 1, 0};
        root = merge(root, tot);
    }
    while(m--){
        int l, r;
        cin >> l >> r;
        auto [L, midr] = split(root, l - 1);
        auto [mid, R] = split(midr, r - l + 1);
        tr[mid].tag ^= 1;
        root = merge(merge(L, mid), R);
    }
    for(int i = 1; i <= n; i++) cout << query(i) << " ";
    return 0;
}

P1486 [NOI2004] 郁闷的出纳员

毕竟有全局加减法,所以我们使用 tag 记录全局加的数。每次进行全局减时,按 k-tag 分裂取 r 即可。

代码如下:

#include <bits/stdc++.h>
using namespace std;
struct treap{
    int val, ls, rs, sz, data, cnt;
}tr[300005];
int n, k, minn, tot, root, tag;
mt19937 rd(time(0));
char op;
pair<int, int> split(int x, int k){
    if(!x) return {0, 0};
    if(tr[x].val < k){
        auto [u, v] = split(tr[x].rs, k);
        tr[x].rs = u;
        tr[x].sz = tr[tr[x].ls].sz + tr[x].cnt + tr[tr[x].rs].sz;
        return {x, v};
    }
    auto [u, v] = split(tr[x].ls, k);
    tr[x].ls = v;
    tr[x].sz = tr[tr[x].ls].sz + tr[x].cnt + tr[tr[x].rs].sz;
    return {u, x};
}
int merge(int l, int r){
    if(!l || !r) return l | r;
    if(tr[l].data < tr[r].data){
        tr[l].rs = merge(tr[l].rs, r);
        tr[l].sz = tr[tr[l].ls].sz + tr[l].cnt + tr[tr[l].rs].sz;
        return l;
    }
    tr[r].ls = merge(l, tr[r].ls);
    tr[r].sz = tr[tr[r].ls].sz + tr[r].cnt + tr[tr[r].rs].sz;
    return r;
}
void insert(int k){
    tr[++tot] = {k, 0, 0, 1, rd(), 1};
    auto [l, r] = split(root, k);
    root = merge(merge(l, tot), r);
}
int find_by_order(int k){
    if(k <= 0) return -1 - tag;
    int pos = root;
    while(pos){
        if(k == tr[tr[pos].ls].sz + 1) return tr[pos].val;
        if(k <= tr[tr[pos].ls].sz) pos = tr[pos].ls;
        else{
            k -= (tr[tr[pos].ls].sz + 1);
            pos = tr[pos].rs;
        }
    }
}
int main(){
    int ans = 0;
    cin >> n >> minn;
    while(n--){
        cin >> op >> k;
        if(op == 'I') if(k >= minn) insert(k - tag);
        if(op == 'A') tag += k;
        if(op == 'S'){
            tag -= k;
            int a = tr[root].sz;
            root = split(root, minn - tag).second;
            ans += a - tr[root].sz;
        }
        if(op == 'F') cout << find_by_order(tr[root].sz - k + 1) + tag << "\n";
    }
    cout << ans;
    return 0;
}

P2042 [NOI2005] 维护数列

这是一道非常复杂的题目,看到区间操作,考虑使用 FHQ-Treap。需要维护区间和、最大前缀和、最大后缀和、最大子段和记作 summaxsum\_lmaxsum\_rmaxsum。需要注意的是,maxsum 的更新比较复杂,转移式为 tr[id].maxsum = max({tr[tr[id].ls].maxsum, tr[tr[id].ls].maxsum_r + tr[id].val, tr[tr[id].ls].maxsum_r + tr[id].val + tr[tr[id].rs].maxsum_l, tr[id].val, tr[id].val + tr[tr[id].rs].maxsum_l, tr[tr[id].rs].maxsum});

然后需要两个懒标记分别为 MAKE-SAME 操作的和 REVERSE 操作的,记作 tag\_make\_sametag\_reversetag\_make\_same 下传优先级更大,下传时应清空 tag\_reverse

参考代码:

#include <bits/stdc++.h>
using namespace std;
int n, m, tot, root;
mt19937 rd(time(0));
stack<int> st;
struct Treap{
    struct node{
        int val, ls, rs, data, sz, tag_make_same = 2000, tag_reverse, sum, maxsum_l, maxsum_r, maxsum;
    }tr[500005];
    void push_down(int id){
        if(tr[id].tag_make_same != 2000){
            if(tr[id].ls){
                tr[tr[id].ls].sum = tr[id].tag_make_same * tr[tr[id].ls].sz;
                tr[tr[id].ls].maxsum = tr[tr[id].ls].maxsum_l = tr[tr[id].ls].maxsum_r = max(tr[id].tag_make_same, tr[id].tag_make_same * tr[tr[id].ls].sz);
                tr[tr[id].ls].tag_make_same = tr[id].tag_make_same;
                tr[tr[id].ls].val = tr[id].tag_make_same;
                tr[tr[id].ls].tag_reverse = 0;              
            }
            if(tr[id].rs){
                tr[tr[id].rs].sum = tr[id].tag_make_same * tr[tr[id].rs].sz;
                tr[tr[id].rs].maxsum = tr[tr[id].rs].maxsum_l = tr[tr[id].rs].maxsum_r = max(tr[id].tag_make_same, tr[id].tag_make_same * tr[tr[id].rs].sz);
                tr[tr[id].rs].tag_make_same = tr[id].tag_make_same;
                tr[tr[id].rs].val = tr[id].tag_make_same;
                tr[tr[id].rs].tag_reverse = 0;              
            }
            tr[id].tag_make_same = 2000;
        }
        if(tr[id].tag_reverse){
            if(tr[id].ls){
                swap(tr[tr[id].ls].maxsum_l, tr[tr[id].ls].maxsum_r);
                swap(tr[tr[id].ls].ls, tr[tr[id].ls].rs);
                tr[tr[id].ls].tag_reverse ^= 1;             
            }
            if(tr[id].rs){
                swap(tr[tr[id].rs].maxsum_l, tr[tr[id].rs].maxsum_r);
                swap(tr[tr[id].rs].ls, tr[tr[id].rs].rs);
                tr[tr[id].rs].tag_reverse ^= 1;             
            }
            tr[id].tag_reverse = 0;
        }
    }
    void push_up(int id){
        tr[id].sz = tr[tr[id].ls].sz + 1 + tr[tr[id].rs].sz;
        tr[id].sum = tr[tr[id].ls].sum + tr[id].val + tr[tr[id].rs].sum;
        if(!tr[id].ls && !tr[id].rs){
            tr[id].maxsum = tr[id].maxsum_l = tr[id].maxsum_r = tr[id].val;
            return;
        }
        if(!tr[id].ls && tr[id].rs){
            tr[id].maxsum = max({tr[id].val, tr[id].val + tr[tr[id].rs].maxsum_l, tr[tr[id].rs].maxsum});
            tr[id].maxsum_l = max(tr[id].val, tr[id].val + tr[tr[id].rs].maxsum_l);
            tr[id].maxsum_r = max(tr[id].val + tr[tr[id].rs].sum, tr[tr[id].rs].maxsum_r);
            return;
        }
        if(tr[id].ls && !tr[id].rs){
            tr[id].maxsum = max({tr[id].val, tr[tr[id].ls].maxsum_r + tr[id].val, tr[tr[id].ls].maxsum});
            tr[id].maxsum_l = max(tr[tr[id].ls].maxsum_l, tr[tr[id].ls].sum + tr[id].val);
            tr[id].maxsum_r = max(tr[tr[id].ls].maxsum_r + tr[id].val, tr[id].val);
            return;
        }
        tr[id].maxsum = max({tr[tr[id].ls].maxsum, tr[tr[id].ls].maxsum_r + tr[id].val, tr[tr[id].ls].maxsum_r + tr[id].val + tr[tr[id].rs].maxsum_l, tr[id].val, tr[id].val + tr[tr[id].rs].maxsum_l, tr[tr[id].rs].maxsum});
        tr[id].maxsum_l = max({tr[tr[id].ls].maxsum_l, tr[tr[id].ls].sum + tr[id].val, tr[tr[id].ls].sum + tr[id].val + tr[tr[id].rs].maxsum_l});
        tr[id].maxsum_r = max({tr[tr[id].rs].maxsum_r, tr[tr[id].rs].sum + tr[id].val, tr[tr[id].rs].sum + tr[id].val + tr[tr[id].ls].maxsum_r});
    }
    pair<int, int> split(int x, int k){
        if(!x) return {0, 0};
        push_down(x);
        if(tr[tr[x].ls].sz < k){
            auto [u, v] = split(tr[x].rs, k - tr[tr[x].ls].sz - 1);
            tr[x].rs = u;
            push_up(x);
            return {x, v};
        }
        auto [u, v] = split(tr[x].ls, k);
        tr[x].ls = v;
        push_up(x);
        return {u, x};
    }
    int merge(int l, int r){
        if(!l || !r) return l | r;
        push_down(l);
        push_down(r);
        if(tr[l].data < tr[r].data){
            tr[l].rs = merge(tr[l].rs, r);
            push_up(l);
            return l;
        }
        tr[r].ls = merge(l, tr[r].ls);
        push_up(r);
        return r;
    }
}tr;
void dfs(int x){
    st.push(x);
    if(tr.tr[x].ls) dfs(tr.tr[x].ls);
    if(tr.tr[x].rs) dfs(tr.tr[x].rs);
}
signed main(){
    int x;
    cin >> n >> m >> x;
    tr.tr[++tot] = {x, 0, 0, (int)rd(), 1, 2000, 0, x, x, x, x};
    root = tot;
    for(int i = 2; i <= n; i++){
        int x;
        cin >> x;
        tr.tr[++tot] = {x, 0, 0, (int)rd(), 1, 2000, 0, x, x, x, x};
        root = tr.merge(root, tot);
    }
    while(m--){
        string s;
        cin >> s;
        if(s == "INSERT"){
            int pos, cnt;
            cin >> pos >> cnt;
            auto [l, r] = tr.split(root, pos);
            for(int i = 1; i <= cnt; i++){
                int x;
                cin >> x;
                if(!st.empty()){
                    tr.tr[st.top()] = {x, 0, 0, (int)rd(), 1, 2000, 0, x, x, x, x};
                    l = tr.merge(l, st.top());
                    st.pop();
                    continue;
                }
                tr.tr[++tot] = {x, 0, 0, (int)rd(), 1, 2000, 0, x, x, x, x};
                l = tr.merge(l, tot);
            }
            root = tr.merge(l, r);
        }
        if(s == "DELETE"){
            int pos, cnt;
            cin >> pos >> cnt;
            auto [l, midr] = tr.split(root, pos - 1);
            auto [mid, r] = tr.split(midr, cnt);
            root = tr.merge(l, r);
            dfs(mid);
        }
        if(s == "MAKE-SAME"){
            int pos, cnt, c;
            cin >> pos >> cnt >> c;
            auto [l, midr] = tr.split(root, pos - 1);
            auto [mid, r] = tr.split(midr, cnt);
            tr.tr[mid].sum = c * tr.tr[mid].sz;
            tr.tr[mid].maxsum = tr.tr[mid].maxsum_l = tr.tr[mid].maxsum_r = max(c, c * tr.tr[mid].sz);
            tr.tr[mid].tag_make_same = c;
            tr.tr[mid].tag_reverse = 0;
            tr.tr[mid].val = c;
            root = tr.merge(tr.merge(l, mid), r);
        }
        if(s == "REVERSE"){
            int pos, cnt;
            cin >> pos >> cnt;
            auto [l, midr] = tr.split(root, pos - 1);
            auto [mid, r] = tr.split(midr, cnt);
            swap(tr.tr[mid].maxsum_l, tr.tr[mid].maxsum_r);
            swap(tr.tr[mid].ls, tr.tr[mid].rs);
            tr.tr[mid].tag_reverse ^= 1;
            root = tr.merge(tr.merge(l, mid), r);
        }
        if(s == "GET-SUM"){
            int pos, cnt;
            cin >> pos >> cnt;
            auto [l, midr] = tr.split(root, pos - 1);
            auto [mid, r] = tr.split(midr, cnt);
            cout << tr.tr[mid].sum << "\n";
            root = tr.merge(tr.merge(l, mid), r);
        }
        if(s == "MAX-SUM") cout << tr.tr[root].maxsum << "\n";
    }
    return 0;
}

练习