AVL平衡树
前言:set 和 map 底层是基于红黑树这颗平衡树来实现的,但由于封装的特性,这颗平衡树的结构对外就是一个黑盒,在算竟中我们往往需要自己维护平衡树结构把树高控制在\(logn\)(可能并不严格),因此需要掌握平衡树的各种操作。
在学习AVL树之前需要对二叉搜索树和平衡搜索树的结构较为熟悉。
AVL 树是一种平衡树,虽然在实践中AVL使用的并不多,但是他对于学习平衡树家族中的其他平衡树具有重要的指导意义,没错现阶段也就意味着进入了平衡树家族的学习阶段,平衡树和线段树一样,代码很长容易写错,在学习的过程(复习)望可以自己手动写一遍无误的代码。
平衡树家族:AVL树,B树,B+树,Treap,FHQ Treap,Splay,替罪羊树,笛卡尔树,K-D树,LCT动态树,红黑树,Size Balance Tree(SBT).
在平衡树家族中,比较常用的是 FHQ Treap,Splay,其中 FHQ Treap 利用动态有序结构的合并和分裂来快速实现弱平衡树结构相关的操作,码量小比较好写,也可以实现持久化;Splay 也是利用合并和分裂来实现平衡树的操作,但是由于其可以快速的实现LCT,也比较常用。对于其他的平衡树等到遇到的时候再做详细记录。
AVL树的整体的结构
AVL树需要维护的信息
需要说明的一点是,在以后的平衡树实现的过程中,不会用 new 等等动态控制堆空间的相关操作,这是效率方面的考究,而是使用动态开点的方式来实现平衡树。
- \(lc_x\):\(x\)节点的左孩子编号,\(rc_x\):\(x\)节点的右孩子编号。
- \(cnt_x\):\(x\)节点出现的次数,\(sz_x\):\(x\)为根的子树中数出现的次数。
- 维护的是可重集,只不过一个节点会记录对应数的出现次数,不会重复插入,后续 FHQ Treap 会把重复的数插入平衡树结构,注意区分。
- \(h_x\):\(x\)节点的高度,注意这里并不是深度。
- \(val_x\):\(x\)节点的权值。
对于这些信息, 可以用结构体维护,但是由于代码里面需要多次访问结构体内部的信息,写起来比较麻烦,所以下面采用多个数组的形式来实现。
左旋和右旋
旋转操作是平衡树 用来维护平衡的重要操作,分为左旋和右旋。
旋转之后节点的信息需要进行更新,更新的方式就类似于线段树节点信息的更新\(pushup\)
code:
void pushup(int x)
{
if(!x) return;
sz[x] = sz[lc[x]] + sz[rc[x]] + cnt[x];
h[x] = max(h[lc[x]], h[rc[x]]) + 1;
}
左旋

- \(y\)为\(x\)节点的右子树,\(x\)节点为需要旋转的节点
- \(x\)节点右子树继承\(y\)节点的左子树,\(rc_x=lc_y\)
- \(y\)节点代替\(x\)节点成为树的根节点,\(x\)节点成为\(y\)节点的左子树,\(lc_y=x\)
code:
void left_rotate(int& x)
{
int y = rc[x];
rc[x] = lc[y];
lc[y] = x;
pushup(x); pushup(y);
x = y;
}
右旋

- \(y\)为\(x\)节点的左子树,\(x\)节点为需要旋转的节点
- \(x\)节点左子树继承\(y\)节点的右子树,\(lc_x=rc_y\)
- \(y\)节点代替\(x\)节点成为树的根节点,\(x\)节点成为\(y\)节点的右子树,\(rc_y=x\)
code:
void right_rotate(int& x)
{
int y = lc[x];
lc[x] = rc[y];
rc[y] = x;
pushup(x); pushup(y);
x = y;
}
旋转调整平衡
即当执行插入和删除操作后,平衡树失衡后如何调整
LL

- 失衡节点左子树高度高于右子树,高度差大于1
- 节点的左子树中,左子树的高度大于等于左子树的右子树的高度
- 右旋\(x\)节点
LR

- 节点左子树的高度高于右子树,高度差大于1
- 节点左子树的左子树高度小于左子树的右子树的高度
- 左旋\(lc_x\),右旋\(x\)
RR

- 失衡节点右子树高度高于左子树高度,高度差大于1
- 右子树的右子树高度大于等于右子树的左子树
- 左旋\(x\)节点
RL

- 失衡节点的右子树高度高于左子树,高度差大于1
- 右子树的右子树高度小于右子树的左子树
- 右旋\(rc_x\),左旋\(x\)。
code:
void rotate(int& x)
{
if(h[lc[x]] - h[rc[x]] > 1) // L
{
if(h[lc[lc[x]]] >= h[lc[rc[x]]]) right_rotate(x); // LL
else
{
// LR
left_rotate(lc[x]);
right_rotate(x);
}
}
else if(h[rc[x]] - h[lc[x]] > 1) // R
{
if(h[rc[rc[x]]] >= h[rc[lc[x]]]) left_rotate(x); // RR
else
{
right_rotate(rc[x]);
left_rotate(x);
}
}
}
插入操作
根据平衡树的规则,节点左子树节点权值比该节点权值小,右子树节点权值比该节点权值大(节点不存储重复出现的元素)。
- \(v=val_x\),$cnt_x++ $
- \(v<val_x\),去往左子树插入
- \(v>val_x\),去往右子树插入
- 如果来到空节点,说明该元素未出现过,为其开辟新编号并记录信息。
code:
void insert(int& x, int v)
{
if(!x)
{
x = ++idx;
val[x] = v;
h[x] = cnt[x] = sz[x] = 1;
return;
}
if(val[x] == v) cnt[x]++;
else if(val[x] > v) insert(lc[x], v);
else insert(rc[x], v);
pushup(x); rotate(x);
}
删除操作
删除操作也比较简单,如果节点计数删除后不为0,则计数--,否则删除该节点。
- \(cnt_x>1\)
- $cnt_x-- $
- \(cnt_x=1\)
- 如果该节点为叶结点或者有为空子树,直接删除或者用子树替代,\(x=lc_x+rc_x\)
- 否则用该节点左子树的最右节点代替该节点,然后继续往下递归删除\(x\)节点,此时情况就会变成第一种直接\(x=lc_x+rc_x\)。
code:
void erase(int& x, int v)
{
if(val[x] == v)
{
if(cnt[x] > 1) cnt[x]--;
else
{
if(!lc[x] || !rc[x]) x = lc[x] + rc[x];
else
{
int y = lc[x];
while(rc[y]) y = rc[y];
val[x] = val[y], cnt[x] = cnt[y];
cnt[y] = 0;
erase(lc[x], val[y]);
}
}
}
else v < val[x] ? erase(lc[x], v) : erase(rc[x], v);
pushup(x); rotate(x);
}
AVL树实现查找操作
查找有多少个元素比当前元素小
- 比当前元素小按照大于等于分组,\(val_x \ge v\),该节点和右子树没有贡献,去左子树找
- \(val_x < v\),累加当前节点计数和左子树计数,去右子树找
code:
int get_rank(int x, int v)
{
if(!x) return 0;
if(val[x] >= v) return get_rank(lc[x], v);
else return sz[lc[x]] + cnt[x] + get_rank(rc[x], v);
}
查找排名为k的元素
- \(sz_{lc_{x}} \ge k\),去左子树找
- \(sz_{lc_x} +cnt_x \ge k\),该节点就是答案
- 否则去右子树找排名第\(k-sz_{lc_x}-cnt_x\)的元素
code:
int get_val(int x, int v)
{
if(sz[lc[x]] >= v) get_val(lc[x], v);
else if(sz[lc[x]] + cnt[x] >= v) return val[x];
else return get_val(rc[x], v - sz[lc[x]] - cnt[x]);
}
查找v的前驱节点
code:
int get_pre(int x, int v)
{
if(!x) return INT_MIN;
if(val[x] >= v) return get_pre(lc[x], v);
else return max(val[x], get_pre(rc[x], v));
}
查找v的后继节点
code:
int get_suf(int x, int v)
{
if(!x) return INT_MAX;
if(val[x] <= v) get_suf(rc[x], v);
else return min(val[x], get_suf(lc[x], v));
}
完整代码
#include <iostream>
using namespace std;
const int N = 1e5 + 10;
int lc[N], rc[N], h[N], sz[N], cnt[N], val[N];
int n, idx, root;
void pushup(int x)
{
if(!x) return;
sz[x] = sz[lc[x]] + sz[rc[x]] + cnt[x];
h[x] = max(h[lc[x]], h[rc[x]]) + 1;
}
void left_rotate(int& x)
{
int y = rc[x];
rc[x] = lc[y];
lc[y] = x;
pushup(x); pushup(y);
x = y;
}
void right_rotate(int& x)
{
int y = lc[x];
lc[x] = rc[y];
rc[y] = x;
pushup(x); pushup(y);
x = y;
}
void rotate(int& x)
{
if(h[lc[x]] - h[rc[x]] > 1) // L
{
if(h[lc[lc[x]]] >= h[lc[rc[x]]]) right_rotate(x); // LL
else
{
// LR
left_rotate(lc[x]);
right_rotate(x);
}
}
else if(h[rc[x]] - h[lc[x]] > 1) // R
{
if(h[rc[rc[x]]] >= h[rc[lc[x]]]) left_rotate(x); // RR
else
{
right_rotate(rc[x]);
left_rotate(x);
}
}
}
void insert(int& x, int v)
{
if(!x)
{
x = ++idx;
val[x] = v;
h[x] = cnt[x] = sz[x] = 1;
return;
}
if(val[x] == v) cnt[x]++;
else if(val[x] > v) insert(lc[x], v);
else insert(rc[x], v);
pushup(x); rotate(x);
}
void erase(int& x, int v)
{
if(val[x] == v)
{
if(cnt[x] > 1) cnt[x]--;
else
{
if(!lc[x] || !rc[x]) x = lc[x] + rc[x];
else
{
int y = lc[x];
while(rc[y]) y = rc[y];
val[x] = val[y], cnt[x] = cnt[y];
cnt[y] = 0;
erase(lc[x], val[y]);
}
}
}
else v < val[x] ? erase(lc[x], v) : erase(rc[x], v);
pushup(x); rotate(x);
}
int get_rank(int x, int v)
{
if(!x) return 0;
if(val[x] >= v) return get_rank(lc[x], v);
else return sz[lc[x]] + cnt[x] + get_rank(rc[x], v);
}
int get_val(int x, int v)
{
if(sz[lc[x]] >= v) get_val(lc[x], v);
else if(sz[lc[x]] + cnt[x] >= v) return val[x];
else return get_val(rc[x], v - sz[lc[x]] - cnt[x]);
}
int get_pre(int x, int v)
{
if(!x) return INT_MIN;
if(val[x] >= v) return get_pre(lc[x], v);
else return max(val[x], get_pre(rc[x], v));
}
int get_suf(int x, int v)
{
if(!x) return INT_MAX;
if(val[x] <= v) get_suf(rc[x], v);
else return min(val[x], get_suf(lc[x], v));
}
int main()
{
cin >> n;
int op, x;
while(n--)
{
cin >> op >> x;
if(op == 1) insert(root, x);
else if(op == 2) erase(root, x);
else if(op == 3) cout << get_rank(root, x) + 1 << endl;
else if(op == 4) cout << get_val(root, x) << endl;
else if(op == 5) cout << get_pre(root, x) << endl;
else cout << get_suf(root, x) << endl;
}
return 0;
}
模板
P3369 【模板】普通平衡树
#include <iostream>
using namespace std;
const int N = 1e5 + 10;
int lc[N], rc[N], h[N], sz[N], cnt[N], val[N];
int n, idx, root;
void pushup(int x)
{
if(!x) return;
sz[x] = sz[lc[x]] + sz[rc[x]] + cnt[x];
h[x] = max(h[lc[x]], h[rc[x]]) + 1;
}
void left_rotate(int& x)
{
int y = rc[x];
rc[x] = lc[y];
lc[y] = x;
pushup(x); pushup(y);
x = y;
}
void right_rotate(int& x)
{
int y = lc[x];
lc[x] = rc[y];
rc[y] = x;
pushup(x); pushup(y);
x = y;
}
void rotate(int& x)
{
if(h[lc[x]] - h[rc[x]] > 1) // L
{
if(h[lc[lc[x]]] >= h[lc[rc[x]]]) right_rotate(x); // LL
else
{
// LR
left_rotate(lc[x]);
right_rotate(x);
}
}
else if(h[rc[x]] - h[lc[x]] > 1) // R
{
if(h[rc[rc[x]]] >= h[rc[lc[x]]]) left_rotate(x); // RR
else
{
right_rotate(rc[x]);
left_rotate(x);
}
}
}
void insert(int& x, int v)
{
if(!x)
{
x = ++idx;
val[x] = v;
h[x] = cnt[x] = sz[x] = 1;
return;
}
if(val[x] == v) cnt[x]++;
else if(val[x] > v) insert(lc[x], v);
else insert(rc[x], v);
pushup(x); rotate(x);
}
void erase(int& x, int v)
{
if(val[x] == v)
{
if(cnt[x] > 1) cnt[x]--;
else
{
if(!lc[x] || !rc[x]) x = lc[x] + rc[x];
else
{
int y = lc[x];
while(rc[y]) y = rc[y];
val[x] = val[y], cnt[x] = cnt[y];
cnt[y] = 0;
erase(lc[x], val[y]);
}
}
}
else v < val[x] ? erase(lc[x], v) : erase(rc[x], v);
pushup(x); rotate(x);
}
int get_rank(int x, int v)
{
if(!x) return 0;
if(val[x] >= v) return get_rank(lc[x], v);
else return sz[lc[x]] + cnt[x] + get_rank(rc[x], v);
}
int get_val(int x, int v)
{
if(sz[lc[x]] >= v) get_val(lc[x], v);
else if(sz[lc[x]] + cnt[x] >= v) return val[x];
else return get_val(rc[x], v - sz[lc[x]] - cnt[x]);
}
int get_pre(int x, int v)
{
if(!x) return INT_MIN;
if(val[x] >= v) return get_pre(lc[x], v);
else return max(val[x], get_pre(rc[x], v));
}
int get_suf(int x, int v)
{
if(!x) return INT_MAX;
if(val[x] <= v) get_suf(rc[x], v);
else return min(val[x], get_suf(lc[x], v));
}
int main()
{
cin >> n;
int op, x;
while(n--)
{
cin >> op >> x;
if(op == 1) insert(root, x);
else if(op == 2) erase(root, x);
else if(op == 3) cout << get_rank(root, x) + 1 << endl;
else if(op == 4) cout << get_val(root, x) << endl;
else if(op == 5) cout << get_pre(root, x) << endl;
else cout << get_suf(root, x) << endl;
}
return 0;
}
P6136 【模板】普通平衡树(数据加强版)
#include <iostream>
#include <climits>
using namespace std;
const int N = 2e6 + 10;
int lc[N], rc[N], cnt[N], h[N], sz[N], val[N];
int n, last, q, idx, root;
void pushup(int x)
{
if(!x) return;
sz[x] = sz[lc[x]] + sz[rc[x]] + cnt[x];
h[x] = max(h[lc[x]], h[rc[x]]) + 1;
}
void left_rotate(int& x)
{
int y = rc[x];
rc[x] = lc[y];
lc[y] = x;
pushup(x); pushup(y);
x = y;
}
void right_rotate(int& x)
{
int y = lc[x];
lc[x] = rc[y];
rc[y] = x;
pushup(x); pushup(y);
x = y;
}
void rotate(int& x)
{
if(h[lc[x]] - h[rc[x]] > 1) // L
{
if(h[lc[lc[x]]] >= h[lc[rc[x]]]) right_rotate(x); // LL
else
{
// LR
left_rotate(lc[x]);
right_rotate(x);
}
}
else if(h[rc[x]] - h[lc[x]] > 1) // R
{
if(h[rc[rc[x]]] >= h[rc[lc[x]]]) left_rotate(x); // RR
else
{
right_rotate(rc[x]);
left_rotate(x);
}
}
}
void insert(int& x, int v)
{
if(!x)
{
idx++;
val[idx] = v;
h[idx] = sz[idx] = cnt[idx] = 1;
x = idx;
return;
}
if(v == val[x]) cnt[x]++;
else if(v < val[x]) insert(lc[x], v);
else insert(rc[x], v);
pushup(x); rotate(x);
}
void erase(int& x, int v)
{
if(!x) return;
if(val[x] == v)
{
if(cnt[x] > 1) cnt[x]--;
else
{
if(!lc[x] || !rc[x]) x = lc[x] + rc[x];
else
{
int y = lc[x];
while(rc[y]) y = rc[y];
cnt[x] = cnt[y]; val[x] = val[y];
cnt[y] = 0;
erase(lc[x], val[y]);
}
}
}
else v < val[x] ? erase(lc[x], v) : erase(rc[x], v);
pushup(x); rotate(x);
}
int get_rank(int x, int v)
{
if(!x) return 0;
if(val[x] >= v) return get_rank(lc[x], v);
else return sz[lc[x]] + cnt[x] + get_rank(rc[x], v);
}
int get_val(int x, int k)
{
if(!x) return 0;
if(sz[lc[x]] >= k) return get_val(lc[x], k);
else if(sz[lc[x]] + cnt[x] >= k) return val[x];
else return get_val(rc[x], k - cnt[x] - sz[lc[x]]);
}
int get_pre(int x, int v)
{
if(!x) return INT_MIN;
if(val[x] >= v) return get_pre(lc[x], v);
else return max(val[x], get_pre(rc[x], v));
}
int get_suf(int x, int v)
{
if(!x) return INT_MAX;
if(val[x] <= v) return get_suf(rc[x], v);
else return min(val[x], get_suf(lc[x], v));
}
void dfs(int root)
{
if(!root) return;
dfs(lc[root]);
cout << val[root] << " " << cnt[root] << endl;
dfs(rc[root]);
}
int main()
{
cin.tie(0); cout.tie(0);
ios::sync_with_stdio(false);
cin >> n >> q;
for(int i = 1; i <= n; i++)
{
int x; cin >> x;
insert(root, x);
}
// dfs(root);
int op, x;
int ans = 0;
while(q--)
{
cin >> op >> x;
x = last ^ x;
// cout << op << " " << x << endl;
// dfs(root);
// cout << endl;
if(op == 1) insert(root, x);
else if(op == 2) erase(root, x);
else if(op == 3) ans ^= (last = (get_rank(root, x) + 1));
else if(op == 4) ans ^= (last = get_val(root, x));
else if(op == 5) ans ^= (last = get_pre(root, x));
else ans ^= (last = get_suf(root, x));
}
cout << ans << endl;
return 0;
}
浙公网安备 33010602011771号