权值线段树

权值线段树的原理与实现:按值域维护元素个数,查询第 K 小与排名。

一句话算法

权值线段树不是维护下标区间,而是维护值域区间;每个叶子表示一个数出现了多少次。

问题模型

普通线段树通常维护数组下标:

index: 1 2 3 4 5 ...
value: a1 a2 a3 ...

权值线段树维护的是数值本身:

value: 1 2 3 4 5 ...
count: 出现次数

如果当前集合是:

3, 3, 5, 8

那么值域上的计数是:

count[3] = 2
count[5] = 1
count[8] = 1

线段树每个节点维护一段值域中出现的总次数。

核心直觉

把每个数字看成一个桶。插入数字 x,就是让 x 这个桶的计数加一;删除数字 x,就是让计数减一。

权值线段树把很多桶按值域区间分组:

[1, 8]
├── [1, 4]
└── [5, 8]

如果左子树一共有 cnt 个数:

  • k 小在左边,当 k <= cnt
  • k 小在右边,当 k > cnt,并且要找右边的第 k-cnt 小。

这就是权值线段树能求第 kk 小的原因。

支持的操作

权值线段树常用来维护一个可重集合,支持:

  1. 插入一个数。
  2. 删除一个数。
  3. 查询某个数的排名。
  4. 查询第 kk 小。
  5. 查询前驱。
  6. 查询后继。

这些操作正好对应普通平衡树模板题。

算法步骤

插入与删除

插入 x

  1. 从根节点表示的整个值域开始。
  2. 根据 x 落在左半还是右半递归。
  3. 沿途所有节点的 sum1
  4. 到达叶子 x

删除 x 同理,只是沿途 sum1

查询排名

x 的排名是集合中小于 x 的数的个数再加一:

rank(x)=count(<x)+1 rank(x)=count(<x)+1

所以只需要查询值域:

[min_value, x - 1]

中的计数和。

查询第 k 小

从根开始:

  1. 设左子树计数为 left_sum
  2. k <= left_sum,答案在左子树。
  3. 否则答案在右子树,令 k -= left_sum
  4. 到叶子时,叶子对应的值就是答案。

查询前驱和后继

前驱是严格小于 x 的最大数:

  1. 先算 cnt = count(<x)
  2. 前驱就是第 cnt 小。

后继是严格大于 x 的最小数:

  1. 先算 cnt = count(<=x)
  2. 后继就是第 cnt+1 小。

算法证明

关键不变量:每个节点的 sum 等于它表示的值域区间中,所有数字的出现次数总和。

插入和删除时,只会影响从根到对应叶子的一条路径。沿途节点的值域都包含这个数,所以它们的 sum 应该同步变化;其它节点不包含这个数,所以不变。不变量保持。

查询第 kk 小时,左子树 sum 表示当前区间左半部分有多少个数:

  • 如果 k 不超过左子树数量,答案一定在左边。
  • 如果 k 超过左子树数量,左边所有数都排在答案前面,答案在右边,并且排名变成 k-left_sum

每次都排除一半值域,直到叶子,得到的值就是第 kk 小。

复杂度分析

设值域大小为 VV

  • 插入:O(logV)O(\log V)
  • 删除:O(logV)O(\log V)
  • 查询排名:O(logV)O(\log V)
  • 查询第 kk 小:O(logV)O(\log V)
  • 查询前驱、后继:O(logV)O(\log V)

如果值域很大但实际出现的数很少,可以使用动态开点,只创建被访问到的节点。

代码实现

cpp
        
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
#include <bits/stdc++.h> using namespace std; struct WeightSegmentTree { struct Node { int left = 0; int right = 0; int sum = 0; }; vector<Node> tree; int root = 0; int min_value; int max_value; WeightSegmentTree(int min_value, int max_value) : min_value(min_value), max_value(max_value) { tree.push_back(Node{}); } int new_node() { tree.push_back(Node{}); return (int)tree.size() - 1; } int add(int u, int l, int r, int pos, int delta) { if (u == 0) u = new_node(); tree[u].sum += delta; if (l == r) return u; int mid = l + (r - l) / 2; if (pos <= mid) { tree[u].left = add(tree[u].left, l, mid, pos, delta); } else { tree[u].right = add(tree[u].right, mid + 1, r, pos, delta); } return u; } void add(int pos, int delta) { root = add(root, min_value, max_value, pos, delta); } int query(int u, int l, int r, int ql, int qr) const { if (u == 0 || qr < l || r < ql) return 0; if (ql <= l && r <= qr) return tree[u].sum; int mid = l + (r - l) / 2; return query(tree[u].left, l, mid, ql, qr) + query(tree[u].right, mid + 1, r, ql, qr); } int count_less(int x) const { if (x <= min_value) return 0; return query(root, min_value, max_value, min_value, x - 1); } int count_leq(int x) const { if (x < min_value) return 0; if (x >= max_value) return tree[root].sum; return query(root, min_value, max_value, min_value, x); } int kth(int u, int l, int r, int k) const { if (l == r) return l; int left_sum = tree[u].left ? tree[tree[u].left].sum : 0; int mid = l + (r - l) / 2; if (k <= left_sum) return kth(tree[u].left, l, mid, k); return kth(tree[u].right, mid + 1, r, k - left_sum); } int kth(int k) const { return kth(root, min_value, max_value, k); } int predecessor(int x) const { int cnt = count_less(x); return kth(cnt); } int successor(int x) const { int cnt = count_leq(x); return kth(cnt + 1); } int size() const { return root == 0 ? 0 : tree[root].sum; } }; int main() { ios::sync_with_stdio(false); cin.tie(nullptr); int n; cin >> n; // Value range for Luogu P3369 ordinary balance tree template. WeightSegmentTree seg(-10000000, 10000000); while (n--) { int op, x; cin >> op >> x; if (op == 1) { seg.add(x, 1); } else if (op == 2) { seg.add(x, -1); } else if (op == 3) { cout << seg.count_less(x) + 1 << '\n'; } else if (op == 4) { cout << seg.kth(x) << '\n'; } else if (op == 5) { cout << seg.predecessor(x) << '\n'; } else if (op == 6) { cout << seg.successor(x) << '\n'; } } return 0; }

测试用例

输入:

10
1 5
1 3
1 7
3 5
4 2
5 5
6 5
2 5
4 2
3 7

输出:

2
5
3
7
7
2

过程:

  • 插入 5,3,7 后,集合为 {3,5,7}
  • 5 的排名是 2
  • 2 小是 5
  • 5 的前驱是 3,后继是 7
  • 删除 5 后,集合为 {3,7},第 2 小是 7

应用分类详解

权值线段树的本质是“在值域上维护频率”。看到动态插入删除数字,同时还要按大小顺序查询时,就应该想到它。

一、普通平衡树操作

典型模式: 动态维护一个可重集合,支持排名、第 kk 小、前驱、后继。

识别信号: 操作包含 insert、delete、rank、kth、predecessor、successor。

核心建模: 每个值的出现次数是叶子计数,节点维护区间总次数。

二、动态中位数

典型模式: 一边插入数字,一边查询当前中位数。

识别信号: 在线维护第 n/2\lceil n/2\rceil 小。

核心建模: 每次插入后查询第 kk 小。

三、逆序对与偏序计数

典型模式: 枚举当前数时,需要知道之前有多少数比它大或比它小。

识别信号: “前面大于它的数”“排名统计”“值域计数”。

核心建模: 插入已扫描数字,查询值域区间计数。

四、值域上的区间统计

典型模式: 问当前集合中落在某个值域范围内的数有多少个。

识别信号: 数值范围查询、频率统计、区间桶。

核心建模: 查询 [L,R] 的计数和。

经典例题

1. 普通平衡树

luogu-P3369

权值线段树可以在值域已知或可动态开点时替代平衡树,完成六种基础操作。

2. 中位数

luogu-P1168

不断插入数字并输出中位数。离散化后用权值线段树维护第 kk 小。

3. 逆序对

从左到右扫描数组。当前数 aia_i 的贡献是已插入数中大于 aia_i 的数量。

参考

  • 本书线段树单点修改:data_structure/segment_tree/update_one/index.md
  • 本书树状数组:data_structure/BIT/index.md