先问一个场景:你们系统里有一批订单金额,随时会有新订单进来,也可能有退款把订单删掉。产品经理每隔几分钟就要看一次“当前金额排第 k 位的订单是多少”。我第一次接到这个需求时,第一反应就是每次现排序。数据量小的时候没感觉,等单量上来,每次 O(n log n) 的排序直接把接口拖慢了一个数量级。后来我换成树状数组维护权值统计,查询的时候在树状数组上做二进制逼近求第 k 小,单次查询 O(log n),代码只有十几行,这一版上线后再也没为这个接口操过心。
这篇文章就专门讲清楚一件事:如何用树状数组求第 k 小的数。我会把原理拆开讲,给出可以直接抄的 C++ 和 Python 模板,再把值域离散化、重复元素、k 的合法性这些常见坑全部列出来。适合正在学算法的同学,也适合在业务里需要维护动态数据流分位数的工程师。看完你不仅能背下模板,还能自己解释清楚为什么这样写是对的。
1. 权值数组:把“找第k小”变成前缀和上的定位
1.1 为什么排序不是万能答案
先想一想,如果数据是静态的,那求第 k 小非常简单:排个序,取下标 k-1 就行。但一旦涉及到“动态”,也就是随时会有元素插入、删除,排序方案就立刻变得很尴尬。每次插入后如果重新排序,复杂度是 O(n log n);如果维护一个有序数组再二分插入,插入本身要移动元素,最坏 O(n)。数据规模到几万、几十万之后,这种做法基本没法用。
这时候需要一个更聪明的思路:我们维护的其实不应该是“序列本身”,而是“值的分布”。什么叫值的分布?就是每个数值出现了多少次。这个结构在算法竞赛里叫权值数组(frequency array)。举个例子,当前集合是 {1, 1, 2, 4},那么:
- 值 1 出现了 2 次,所以
freq[1] = 2 - 值 2 出现了 1 次,所以
freq[2] = 1 - 值 3 出现了 0 次,所以
freq[3] = 0 - 值 4 出现了 1 次,所以
freq[4] = 1
有了这个频次数组,插入一个元素 x 就是 freq[x]++,删除一个元素 x 就是 freq[x]--,单次更新 O(1),看起来很不错。但问题来了:怎么求第 k 小?
1.2 权值数组与前缀和的单调性
定义前缀和 S(i) = freq[1] + freq[2] + ... + freq[i],它表示“值小于等于 i 的元素总个数”。这个函数有一个非常漂亮的性质:它是单调不减的。因为 freq 数组每一项都大于等于 0,往右累加只会越来越多。
利用这个单调性,第 k 小就可以翻译成一句话:找到最小的下标 i,使得 S(i) >= k。
用刚才的例子验证一下。集合 {1, 1, 2, 4},前缀和是:
| i | 1 | 2 | 3 | 4 |
|---|---|---|---|---|
| freq[i] | 2 | 1 | 0 | 1 |
| S(i) | 2 | 3 | 3 | 4 |
现在问第 3 小是多少。S(1) = 2 < 3,不行;S(2) = 3 >= 3,所以最小的满足条件的是 i=2。第 3 小就是 2。手动排序验证一下:{1, 1, 2, 4} 排序后第 3 个确实是 2,完全正确。
所以“求第 k 小”本质上是一个“在单调序列上做定位”的问题,和二分查找是天生一对。
1.3 树状数组在其中扮演的角色
有了权值数组和前缀和的想法,剩下的就是找一个数据结构,让它既能快速更新某个位置的 freq,又能快速查询前缀和。
如果用普通数组,更新 O(1),但查询某个前缀和需要从 1 累加到 i,最坏 O(n)。如果查询频率高,这个代价还是扛不住。
树状数组(Binary Indexed Tree,简称 BIT)就是专门干这个的。它用 lowbit 把数组拆成若干段,每个节点 c[i] 保存的是一个区间 [i - lowbit(i) + 1, i] 的和。单点更新和前缀和查询都能做到 O(log n)。
到这里,方案已经呼之欲出了:用树状数组维护 freq 数组,查询第 k 小时在值域上二分,每次通过 query(mid) 判断是否满足条件。这个做法复杂度是 O(log n * log n),已经可以解决不少问题。但既然上了树状数组,其实还有更优的玩法——直接在树状数组上二分,把复杂度压到 O(log n)。这就是下一章的主角。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 二进制位逼近:树状数组上二分到底在做什么
2.1 朴素二分的复杂度瓶颈
很多文章和博客会给出这样一段代码:在值域 [1, n] 上二分 mid,每次调用 sum(mid) 判断 sum(mid) >= k 是否成立。如果成立,就把右边界收缩到 mid;否则把左边界移动到 mid+1。
这个做法思路清晰,完全没错,但复杂度是 O(log n * log n)。外层二分是 O(log n),每次内层 query 又是 O(log n)。当查询次数是 10 万、值域是 10 万时,这个常数实际上是可以接受的。但如果查询次数到 100 万甚至更高,或者值域到 10 的 9 次方(离散化后还是 10 万),你会希望把查询压到单次 O(log n)。
树状数组恰好支持这一点,因为它的 c[i] 存的是某一段区间的和。我们不是去二分整个值域,而是直接在树状数组的节点上“跳”,一次跳一个二进制位,全程只查 c[i],不再调 query。
2.2 c[nxt] 为什么恰好是一段累加和
先回顾树状数组的定义:c[i] 保存 freq[i - lowbit(i) + 1] 到 freq[i] 的和。也就是说,c[i] 覆盖的是一个长度为 lowbit(i) 的连续区间,区间的右端点是 i。
现在考虑我们在做什么。假设 pos 是当前已经确认“前缀和仍然小于 k”的最大下标,初始为 0。我们从大到小枚举一个二进制位 step = 2^p,看看能不能跳到 nxt = pos + step。
这里有个关键点:因为我们是从高位到低位逐个尝试的,所以当处理到第 p 位时,pos 的二进制表示里,第 p 位以及更低位都是 0。为什么?因为更低位还没被处理过,而当前位也还没被决定。于是 pos + step 这个操作不会产生进位,nxt 的 lowbit 恰好等于 step。
这个性质非常关键,它意味着 c[nxt] 正好覆盖区间 [pos + 1, pos + step],不多不少。也就是说,我可以用一次 O(1) 的数组访问,拿到从“当前已确定位置的下一个位置”开始、长度为 step 的这一段的和。
如果你理解了这一点,就已经理解了树状数组上二分的 80%。剩下的只是循环条件怎么写的问题。
2.3 标准实现与手工推演
先给 C++ 实现,后面再给 Python 版本。这里 n 是树状数组的大小,即最大权值下标;bit 是树状数组本体。假设元素已经通过 add 插入到树状数组里了:
cpp复制// 求第 k 小,返回权值下标,下标范围 [1, n]
int kth(int k) {
int pos = 0;
int cur = 0; // 当前已累加的前缀和
int step = 1;
while ((step << 1) <= n) step <<= 1; // 找到不超过 n 的最大 2 的幂
for (; step; step >>= 1) {
int nxt = pos + step;
if (nxt <= n && cur + bit[nxt] < k) {
cur += bit[nxt];
pos = nxt;
}
}
return pos + 1;
}
可能光看代码还是有点抽象,我们拿一个具体例子完整走一遍。假设当前集合是 A = {3, 1, 4, 1, 5, 9, 2, 6},值域最大 9。先统计 freq:
| 值 | 1 | 2 | 3 | 4 | 5 | 6 | 7 | 8 | 9 |
|---|---|---|---|---|---|---|---|---|---|
| freq | 2 | 1 | 1 | 1 | 1 | 1 | 0 | 0 | 1 |
构建树状数组后,关键位置的 bit 值为:bit[1]=2,bit[2]=3,bit[3]=1,bit[4]=5,bit[6]=2,bit[8]=7,bit[9]=1。现在求第 4 小。n=9,初始 step=8,pos=0,cur=0。
| step | nxt | cur + bit[nxt] | 是否 < 4 | 操作 |
|---|---|---|---|---|
| 8 | 8 | 0 + 7 = 7 | 否 | 跳过 |
| 4 | 4 | 0 + 5 = 5 | 否 | 跳过 |
| 2 | 2 | 0 + 3 = 3 | 是 | pos=2, cur=3 |
| 1 | 3 | 3 + 1 = 4 | 否 | 跳过 |
最后返回 pos + 1 = 3。对一下排序结果 {1, 1, 2, 3, 4, 5, 6, 9},第 4 个确实是 3。整个过程中一次 query 都没调,只访问了 bit[nxt],所以单次查询就是 O(log n)。
2.4 为什么从大到小枚举 step 不会重复也不会漏
这个问题的本质是:我们要找的答案 ans,它的二进制表示本身就是由若干 2 的幂拼出来的。算法从最高位开始,逐个确定 ans 的每一位是 0 还是 1。
判断方法很朴素:如果当前已经累加的和,加上这一位覆盖的那一段连续区间的和,仍然小于 k,那么第 k 小一定还在更右边,这位可以取 1;否则这位只能取 0,继续看更低位。
因为 pos 在被某个 step 更新后,它的低位全部是 0(这是从高位到低位处理的必然结果),所以下一次取 bit[pos + step] 时,覆盖的区间一定是从 pos + 1 开始的紧挨着的下一段,不会和之前累加过的区间重叠,也不会漏掉中间任何一段。整个过程相当于用二进制“装配”出 ans - 1,而 ans 就是第一个让前缀和大于等于 k 的位置。
这个性质用一句话记忆就是:树状数组节点天然是“按 lowbit 划分的连续区间块”,二进制从大到小枚举,正好把这些块按顺序拼起来。
3. 完整模板与四个高频踩坑点
3.1 C++ / Python 可直接抄的模板
为了方便直接使用,我把完整版的 C++ 模板放在这里。包含插入、单点删除、前缀和查询、求第 k 小四个基础操作。
cpp复制const int MAXN = 100000 + 10;
int n; // 值域大小
long long bit[MAXN]; // 树状数组本体,用 long long 更稳
void add(int idx, long long delta) {
for (; idx <= n; idx += idx & -idx) {
bit[idx] += delta;
}
}
long long preSum(int idx) {
long long res = 0;
for (; idx > 0; idx -= idx & -idx) {
res += bit[idx];
}
return res;
}
int kth(long long k) {
int pos = 0;
long long cur = 0;
int step = 1;
while ((step << 1) <= n) step <<= 1;
for (; step; step >>= 1) {
int nxt = pos + step;
if (nxt <= n && cur + bit[nxt] < k) {
cur += bit[nxt];
pos = nxt;
}
}
return pos + 1;
}
Python 版本同样很简洁:
python复制class BIT:
def __init__(self, n):
self.n = n
self.bit = [0] * (n + 1)
def add(self, idx, delta):
while idx <= self.n:
self.bit[idx] += delta
idx += idx & -idx
def prefix(self, idx):
res = 0
while idx > 0:
res += self.bit[idx]
idx -= idx & -idx
return res
def kth(self, k):
pos = 0
cur = 0
step = 1 << (self.n.bit_length() - 1)
while step:
nxt = pos + step
if nxt <= self.n and cur + self.bit[nxt] < k:
cur += self.bit[nxt]
pos = nxt
step >>= 1
return pos + 1
两个版本逻辑完全一致。使用的时候注意,kth 返回的是权值下标,不是原值。如果你做过离散化,拿到下标后还要映射回真实值。
3.2 坑一:k 的合法性和空集合
这是一个非常容易被忽略的问题。kth 函数内部没有检查 k 是否合法,如果 k <= 0,那循环条件 cur + bit[nxt] < k 永远为真,最后会一直跳到 pos = n,返回 n + 1。如果 k > total(当前集合元素总数),算法会把所有元素都累加进 cur,同样返回 n + 1。
所以调用前一定要保证:
cpp复制if (total > 0 && k >= 1 && k <= total) {
int ans = bit.kth(k);
}
空集合时 total == 0,此时不存在任何合法的第 k 小,必须先判空。这个边界在写动态中位数、动态排行榜时特别容易踩,因为数据流是慢慢来的,一开始就可能是空的。
另外要记得,树状数组的下标是从 1 开始。如果你统计的值域是从 0 开始的,插入时要统一 idx + 1,否则 lowbit 的过程会在下标 0 处死循环。
3.3 坑二:重复元素与“第 k 个不同的数”不是一回事
freq[x] 记录的是值 x 出现的次数。用这个结构求出的第 k 小,是“排序后第 k 个位置上的值”。也就是说,重复元素会占多个位置。
举个例子:集合 {1, 1, 2, 4},第 2 小是 1,第 3 小是 2。这个语义符合大多数“第 k 小”的题目定义,没有问题。
但如果你要的是“第 k 个不同的值”,比如集合同样是 {1, 1, 2, 4},你希望第一个不同值是 1,第二个不同值是 2,第三个不同值是 4。那么上面的模板就不适用了,因为 freq 里的累计值是 2、1、1,而不是 1、1、1。
处理方式有两种:一是把所有非零频次改成 1 再查询;二是维护一个“去重后的值列表”,在这个列表上用不同结构查询。具体用哪种取决于业务。我的建议是,在动手之前先把需求语义确认清楚,这是我被坑过最多次的地方。
3.4 坑三:值域太大必须离散化
如果值域是 [1, 10^9],你不可能开一个长度为 10 亿的数组。这时候需要离散化,把所有出现过的值收集起来,排序去重,映射成 1..m,其中 m 是不同值的个数。BIT 的大小是 m,不是原值域。
离散化之后,kth 返回的是压缩后的下标 id,真实值是去重数组里的第 id 个元素。这个映射关系在写代码时很容易搞丢,我建议封装一个小函数统一处理:
cpp复制vector<long long> vals; // 排序去重后的所有可能值
// 查询第 k 小对应的真实值
long long getRealValue(int kthResult) {
return vals
