ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

快速选择算法:第k小元素、O(n)分区与TopK工程实践

快速选择算法:第k小元素、O(n)分区与TopK工程实践 1. 一道面试题背后的真实需求第k小元素到底在哪些场景里冒出来给定一个无序序列找出其中第 k 小的元素。这句话出现在无数算法题里很多人第一反应是排序然后取arr[k-1]。这个答案在面试中通常会被追问一句能不能更快如果你只是把它当成刷题任务很可能就此止步但如果把它当成工程问题你会发现它几乎贯穿了数据处理中的一大类需求。我第一次认真对待这个问题是在做一个后台统计模块的时候。需求描述得非常朴素给一批用户的行为分数据找出中位数和 95 分位数用于判断整体水位是否异常。数据量是千万级每隔几分钟刷新一次。当时的实现就是把整批数据排一遍序再按下标取。单机跑一次要好几秒CPU 全程打满而这批数据的绝大部分排列顺序对最终结果毫无意义——我只要两个位置上的值却为全部位置都买了单。这就是选择问题和排序问题的分界线。选择问题Selection Problem指的是在一个包含 n 个元素的集合中找到第 k 个顺序统计量。k1 是最小值kn 是最大值k(n1)/2 是中位数。它和排序的关系很微妙——排序顺带解决了选择但选择并不需要排序的全部信息。这个不需要三个字就是分治思想能在这里大展拳脚的原因。这篇文章适合谁看如果你正在学算法设计与分析想真正搞懂分治法怎么落地那本文会把分区的每一步拆开讲如果你是有一定经验的开发者知道快速排序但没认真写过快速选择那本文会告诉你 pivot 选不好会退化到什么程度、重复元素怎么处理、k 的下标从 0 还是 1 开始这些实际会踩的坑如果你只想快速拿到一份能跑的代码第 4 节的实现可以直接抄。我尽量不写空话每个决策背后的为什么都会交代清楚。2. 先把账算清楚O(n log n)和O(n)之间的差距有多大2.1 排序的开销到底花在哪里比较排序有一个理论下界O(n log n)。这不是工程实现的疏忽而是信息论决定的——n 个元素有 n! 种排列可能每次比较最多区分两种结果所以要区分全部排列至少需要 log₂(n!) 次比较而 log₂(n!) 约等于 n log n。这个下界对排序问题成立但注意它对选择问题并不成立。原因很直白排序要求输出一个完整的有序序列而选择只要求你报出第 k 个位置是哪个值。后者需要确定的信息量小得多所以下界也低。事实上第 k 小元素可以在O(n)时间内确定而且这个 O(n) 不是平均意义下的 O(n)那么简单——有一类算法能做到最坏情况 O(n)。打个比方你要在一场马拉松里找出第 100 名选手。排序做法是让全部选手按到达顺序排好队再数到第 100 个而选择做法只需要知道有多少人比某个成绩快然后不断缩小范围根本不需要给每个人排出精确名次。2.2 选择问题为什么可以低于排序下界关键动作叫partition分区随便挑一个元素当基准pivot把序列重新排列成左边都比它小右边都比它大的两堆。这一步只花 O(n) 时间一次线性扫描就够了。做完之后pivot 的最终位置就确定了假设它落在下标 p 上。接下来就是分治的剪枝如果 p 正好等于我们要找的 k答案就是它如果 p 比 k 大说明目标在左半边右边那一堆可以整个丢掉如果 p 比 k 小目标在右半边左边那一堆丢掉。每一轮只递归一个分支这就是它区别于快速排序两个分支都要处理的核心。快速排序每层处理 n 个元素、共 log n 层总代价 n log n快速选择每层处理规模递减的一段期望总代价是 n n/2 n/4 ... ≈ 2n也就是 O(n)。2.3 一个具体的数据对比光说复杂度不够直观我拿实际跑过的数字做个对比。测试机是一台普通的开发机数据是随机生成的 32 位整数重复跑 5 次取平均数据规模 n全排序耗时快速选择耗时中位数查询耗时10 万约 9 ms约 1.2 ms约 1.1 ms100 万约 110 ms约 12 ms约 11 ms1000 万约 1.4 s约 130 ms约 125 ms表里的数字只是量级参考不同语言、不同编译器会有出入但趋势很稳定规模越大差距越明显。排序耗时随 n 的增长快于线性而快速选择基本贴着线性走。当 n 从 100 万涨到 1000 万排序慢了 10 倍多快选只慢了约 10 倍两个都接近线性这里要小心1.4s 对 110ms 是 12.7 倍而 130ms 对 12ms 是 10.8 倍差距在扩大只是在这个数据范围内还不够夸张。如果把缓存局部性、内存带宽这些因素算进去n 再大一个量级差距会进一步拉开。3. QuickSelect的骨架partition是整个算法的发动机3.1 Lomuto分区最好写好懂的版本先看最经典、最容易写对的 Lomuto 分区方案。它的思路是选最右边的元素当 pivot然后用一个指针i标记小于区的右边界另一个指针j从左往右扫。遇到比 pivot 小的元素就把它换到i的位置i右移。int partitionLomuto(vectorint a, int lo, int hi) { int pivot a[hi]; // 取最右元素作基准 int i lo; // i 指向小于区的下一个空位 for (int j lo; j hi; j) { if (a[j] pivot) { swap(a[i], a[j]); i; } } swap(a[i], a[hi]); // 把基准放到它最终该在的位置 return i; // 返回基准的最终下标 }这段代码有几个细节值得抠。第一循环只扫到hi - 1因为a[hi]是基准本身不需要和自己比较。第二最后的一次swap(a[i], a[hi])是把基准放到分界点上这一步不能少否则基准还在最右边左边的小于区就失去了分界意义。第三返回的i就是基准的最终下标这个下标在后续所有轮次里都不会再变。Lomuto 的优点是直观缺点是交换次数偏多即使a[j]已经在自己该在的位置上也会做一次自我交换。对大规模数据来说这是个常数级的损耗不算致命但如果你的场景对写操作敏感比如元素是很大的结构体swap 很贵就得考虑下一个方案。3.2 Hoare分区更少的交换次数Hoare 分区用双指针从两端向中间逼近。左指针找第一个不小于 pivot 的元素右指针找第一个不大于 pivot 的元素然后交换直到两指针相遇。它和 Lomuto 最大的区别在于返回值不是基准的下标而是左右两个子区间的分界点而且分界点两侧的区间可能和基准有重叠。int partitionHoare(vectorint a, int lo, int hi) { int pivot a[lo (hi - lo) / 2]; // 取中间元素的值作基准 int i lo - 1, j hi 1; while (true) { do { i; } while (a[i] pivot); do { --j; } while (a[j] pivot); if (i j) return j; swap(a[i], a[j]); } }Hoare 版本平均交换次数大约是 Lomuto 的三分之一很多标准库的排序实现内部都用类似思路。但它有个容易翻车的点返回的是 j 不是 i而且递归调用时右半边的起点是j 1不是j。写错一个下标轻则死循环重则越界。如果你对边界不自信建议先用 Lomuto 版本跑通逻辑再考虑换成 Hoare 优化。3.3 递归只有一个分支这就是它能到O(n)的原因把分区接上选择逻辑整个骨架就成型了。下面是迭代写法避免了递归调用栈的额外开销int quickSelect(vectorint a, int k) { // k 为 0-based 下标找第 k1 小元素 int lo 0, hi (int)a.size() - 1; while (true) { if (lo hi) return a[lo]; // 区间只剩一个元素 int p partitionLomuto(a, lo, hi); if (p k) return a[p]; // 命中 if (k p) hi p - 1; // 目标在左半边 else lo p 1; // 目标在右半边 } }请特别注意while(true)里那两次区间收缩。它和快速排序的递归结构长得像但本质完全不同快排会对左右两半都发起递归所以递归树有两个分支、深度 log n快速选择只走进其中一个分支所以每层只处理一个子问题递归深度在期望意义下是 O(log n)但最坏情况会退化到 O(n) 层。这就是为什么我把循环写成迭代形式。虽然期望递归深度不高但最坏情况下递归版本可能压爆栈迭代版本则完全没有这个顾虑同时也省掉了函数调用开销。你可以把lo hi这个判断理解成区间收缩到只剩一个元素它必然就是答案这是循环能正确终止的最后一道保险。4. 从伪代码到能跑的程序逐行拆解与实测4.1 C版本完整实现与随机化改造上一节的版本有个隐患固定取最右元素作 pivot在近乎有序的输入上会退化。你可能会说第 k 小元素问题输入是无序序列怕什么有序现实中的输入往往没那么随机——日志按时间写入、ID 单调递增、某个字段被批量更新过都可能让数据呈现局部有序。所以工程实现里几乎都会加一步随机化#include bits/stdc.h using namespace std; static mt19937 rng(chrono::steady_clock::now().time_since_epoch().count()); int partitionRandom(vectorint a, int lo, int hi) { uniform_int_distributionint dist(lo, hi); int idx dist(rng); swap(a[idx], a[hi]); // 随机挑一个换到最右 int pivot a[hi]; int i lo; for (int j lo; j hi; j) { if (a[j] pivot) { swap(a[i], a[j]); i; } } swap(a[i], a[hi]); return i; } int quickSelect(vectorint a, int k, int lo, int hi) { while (true) { if (lo hi) return a[lo]; int p partitionRandom(a, lo, hi); if (p k) return a[p]; if (k p) hi p - 1; else lo p 1; } } int main() { vectorint a {7, 1, 9, 3, 5, 2, 8, 6, 4, 0}; int k 4; // 0-based即第 5 小 cout quickSelect(a, k, 0, (int)a.size() - 1) endl; // 输出 4 return 0; }这里的rng用了mt19937而不是老旧的rand()。原因是rand()在很多平台上范围只有 32767且分布质量差用它取模会引入明显的偏置。mt19937的周期是 2^19937-1质量高得多代价是每个随机数约几十纳秒对分区这种量级的操作完全可以接受。还有一点随机数种子不要固定成常数。如果你在性能测试里固定种子测出来的结果可能因为某次特别坏的 pivot 选择而剧烈波动。用chrono::steady_clock取种子每次运行都不一样测出来的才是分布意义上的真实表现。4.2 Python版本的递归深度问题Python 写起来更短但有两个坑必须提前说。第一递归版本的快速选择在最坏情况下深度可能达到 n而 Python 默认递归上限是 1000n 稍大就会RecursionError第二Python 每次函数调用开销大递归写法在千万级数据上性能不如迭代。import random def quick_select(a, k): a: 可变序列k: 0-based 下标返回第 k1 小元素 lo, hi 0, len(a) - 1 while True: if lo hi: return a[lo] p partition(a, lo, hi) if p k: return a[p] elif k p: hi p - 1 else: lo p 1 def partition(a, lo, hi): idx random.randint(lo, hi) a[idx], a[hi] a[hi], a[idx] pivot a[hi] i lo for j in range(lo, hi): if a[j] pivot: a[i], a[j] a[j], a[i] i 1 a[i], a[hi] a[hi], a[i] return i迭代版本彻底绕开了递归栈的问题。我实测过在 CPython 3.11 上跑 100 万个随机整数迭代版约 0.6 秒递归版约 0.9 秒差距主要来自函数调用。如果你用的是 PyPy差距会缩小但迭代版依旧更稳。注意Python 里a[x], a[y] a[y], a[x]是原子操作不会出现中间状态但它是先构造元组再解包比手动中转变量略慢。对性能极其敏感的场景可以试试tmp a[i]; a[i] a[j]; a[j] tmp实测有微弱的提升。4.3 一遍跑通的验证方法对拍写完算法最怕什么怕逻辑错但恰好在你的测试用例上过了。我的习惯是写一个对拍脚本用暴力解法当参照import random def brute_force(a, k): return sorted(a)[k] random.seed(42) for case in range(2000): n random.randint(1, 60) arr [random.randint(-20, 20) for _ in range(n)] k random.randrange(n) got quick_select(arr[:], k) want brute_force(arr[:], k) assert got want, fcase{case} arr{arr} k{k} got{got} want{want} print(all passed)这个脚本有三个设计点。第一数值范围故意设得很小-20 到 20这样能制造大量重复元素专门用来暴露没有正确处理等于 pivot 的元素这类 bug。第二题量 2000 组覆盖了各种 n 和 k 的组合比手写几个用例靠谱得多。第三传参用arr[:]拷贝防止第一遍执行把数组改乱影响第二遍的参照结果。我自己的经验是快速选择的所有 bug 几乎都能被这个对拍脚本抓到。抓不到的通常是性能问题比如在大量重复元素上退化那需要另外构造测试用例第 6 节会讲。5. pivot选不好真的会退化最坏情况的触发条件与三种应对5.1 什么时候会踩到O(n²)快速选择最坏情况是 O(n²)触发条件非常明确每一轮分区后pivot 都恰好落在当前区间的某一端也就是分区结果一边是 0 个元素、另一边是 n-1 个元素。此时每轮只减少一个规模总代价是 n (n-1) ... 1 n(n1)/2。什么时候会出现这种极端情况最典型的是固定 pivot 输入近乎有序。假设你取最右元素作 pivot输入是严格递增的[1,2,3,...,n]那么第一轮 pivot 是 n它本来就是最大值分区后左边 n-1 个元素、右边 0 个递归到左边pivot 又是 n-1 的最大值……于是完美退化成 O(n²)。对于已经排序或者逆序的数据这个行为是确定性的不是概率问题。这也是我强调必须做随机化的原因随机化把输入对抗算法变成了算法自身随机。即使输入本身有序随机选中的 pivot 期望落在中间附近每轮期望砍掉一半。5.2 随机化pivot期望复杂度的由来随机化的复杂度分析可以这样理解。假设当前区间长度是 mpivot 随机落在任意位置的概率相同。那么有 50% 的概率pivot 落在中间一半的范围里此时两个子区间里较大的那个长度不超过 3m/4。也就是说每两轮有效分区至少把规模缩小到原来的 3/4。所以规模的期望收缩序列是 m, (3/4)m, (3/4)²m, ...直到收缩到 1迭代次数约是 log_{4/3} m对应到总代价上每层期望处理量不超过 m于是总期望代价是对一个几何级数求和结果是 O(m)。严格证明会用递归式或期望分析但上面这个每两轮砍掉四分之一的直觉已经足够解释为什么它能稳定在线性附近。需要强调的是随机化的 O(n) 是期望意义上的不是最坏意义上的。理论上存在极小的概率每次随机都选中端点的元素算法仍然退化。概率低到可以忽略——每次选中端点的概率约是 2/m连续多次命中端点的概率是阶乘级的下降——但可以忽略和不存在是两码事对实时性要求极高的系统这一点要心里有数。5.3 三数取中与中位数的中位数BFPRT如果你不满足于期望 O(n)想要最坏 O(n)那就得上 BFPRT 算法也叫中位数的中位数算法。它的核心思路是把大问题拆成确定性的小问题把 n 个元素分成 ⌈n/5⌉ 组每组 5 个元素最后一组可能不足 5 个找出每一组的中位数得到一个长度为 ⌈n/5⌉ 的中位数序列递归调用本算法找出这个中位数序列的中位数作为 pivot用这个 pivot 做分区然后和快速选择一样只递归一个分支。为什么这样做能保证最坏 O(n)关键在于这个 pivot 的性质它至少比 3n/10 个元素大也至少比 3n/10 个元素小。因为每组 5 个里有 3 个中位数及其两侧能被确定地划分到 pivot 的某一侧而大约一半的组的中位数小于 pivot。这样一次分区至少能砍掉 30% 的规模规模收缩就变成了几何级数最坏也是 O(n)。代价是常数非常大分组、求组内中位数5 个元素的小排序可以直接用插入排序、递归求中位数这些操作加起来让它在实际运行中往往是普通随机化快速选择的数倍慢。所以我个人的选择是工程里默认用随机化快速选择只有在做算法研究或者对最坏情况有硬性要求的场景才用 BFPRT。三数取中取首、中、尾三个数的中位数作 pivot是个折中方案对付部分有序数据效果不错但它给出的依然是期望复杂度不能提供最坏保证。5.4 什么时候该用堆而不是快选还有个常见误区以为找第 k 小用快选永远最优。其实要看你需要什么需求推荐方案理由单次查询数据在内存随机化快速选择期望 O(n)常数小多次查询同一数组的任意 k先排序 / 建索引一次 O(n log n) 换后续 O(1)求 TopKk 远小于 n大顶堆维护 k 个O(n log k)不用改原数组数据流k 固定且只增不减双堆结构增量更新插入 O(log n)要求最坏 O(n)BFPRT提供确定性最坏保证表里TopK 用堆这一行特别值得说。当 k 很小比如只要前 10 名时维护一个大小为 k 的大顶堆遍历一遍数组元素比堆顶小就替换堆顶总代价 O(n log k)。它和快速选择都是 O(n) 级别但堆方案不需要修改原数组这在原数据不可写或者需要并发读取的场景下是决定性的优势。而且堆方案的复杂度是稳定 O(n log k)没有随机化那点不确定性。6. 那些反复调不对的边界下标、重复元素与k的起点6.1 k从0还是1开始这是新手最容易栽的第一个跟头。数学表述里第 k 小元素通常是 1-basedk1 是最小值而数组下标是 0-based。写代码时如果不统一就会出现差一错误。我的建议是函数内部统一用 0-based也就是找下标为 k 的元素然后在文档注释里明确写清楚k0 返回最小值kn-1 返回最大值。调用方如果想找第 5 小传k4。我在项目里见过更隐蔽的写法函数参数叫k内部用k - 1转换但某次重构时把转换去掉了测试用例又恰好都用了 k1结果全过。等上线后很多人反馈结果偏差一个位置排查了很久。教训是下标约定要么写进注释要么写进函数名比如quick_select_zero_based这种命名虽然啰嗦但能救命。6.2 大量重复元素时的三路划分第二个坑更隐蔽当数组中大量元素相等时Lomuto 分区的行为会变差。考虑一个极端例子——数组里 100 万个元素全是 5。用 Lomuto 分区a[j] pivot永远为假所以所有元素都被划到大于等于区pivot 落在最左边下一轮区间只缩小 1。结果就是 O(n²)而且这个退化不是随机的是必然的随机化也救不了因为无论选哪个元素作 pivot值都是 5。解决方案是三路划分也叫荷兰国旗问题一次分区把数组分成三段小于 pivot、等于 pivot、大于 pivot。def partition3(a, lo, hi): 返回 (lt, gt)[lo, lt) 全小于pivot(gt, hi] 全大于pivot[lt, gt] 全等于pivot import random pivot a[random.randint(lo, hi)] lt, i, gt lo, lo, hi while i gt: if a[i] pivot: a[lt], a[i] a[i], a[lt] lt 1; i 1 elif a[i] pivot: a[gt], a[i] a[i], a[gt] gt - 1 # 注意这里不自增 i换过来的元素还没检查 else: i 1 return lt, gt def quick_select_dup(a, k): lo, hi 0, len(a) - 1 while True: if lo hi: return a[lo] lt, gt partition3(a, lo, hi) if lt k gt: return a[k] # 目标落在等于区直接命中 elif k lt: hi lt - 1 else: lo gt 1三路划分的关键在于等于区是一次性归位的。如果 k 落在等于区里直接返回就行因为这一段里所有值都相同位置随便取一个都是正确答案。对于全相同的数组第一次分区后ltlo、gthik 一定落在等于区内一轮结束复杂度 O(n)。这里还有个细节a[i] pivot分支里交换完之后i不能自增。原因是从gt换过来的那个元素还没有检查过必须留在原地等下一轮循环判断。这个点我第一次写的时候也搞错了表现是部分元素被跳过结果时对时错。6.3 递归终止条件的几个易错写法我把踩过的终止条件问题整理成一张表都是自己或同事真实写错过的错误写法后果正确做法只判断lo hi但用h lo循环空区间时越界访问先判断lo hi或保证循环内区间非空收缩时写hi p而不是hi p - 1若 pk 未命中区间不收缩死循环命中情况单独返回未命中必须排除 pHoare 分区后用lo j左右区间重叠可能死循环Hoare 的右半从j 1开始三路划分后k落在等于区却继续递归无意义的递归甚至栈溢出等于区直接返回其中死循环是最难查的因为程序不报错就是一直跑。排查方法很简单在循环里打印每次的lo、hi、p看区间有没有真的在收缩。如果某次迭代前后区间完全一样问题就出在收缩逻辑上。我个人的习惯是在写循环版本时先确认三件事循环出口在哪、区间是否严格收缩、k 的边界是否可能越界。这三件事想清楚了代码基本不会错。提示如果调用方传进来的 k 可能越界比如 k n务必在函数入口做一次检查。我在一个线上模块里就遇到过 k 传成 100 而数组只有 50 个元素的情况因为没做校验循环一路收缩到越界返回了一个垃圾值而且不报错问题很难被发现。7. 把选择算法放进更大的场景里从静态数组到数据流7.1 静态数组重复查询先排序的取舍前面一直在说快选比排序快但这有个前提只查一次。如果你的业务需要反复查询同一个数组的不同 k那结论可能反过来。假设你要查询 100 次每次 k 都不同每次都跑快速选择100 × O(n) O(100n)先排序一次O(n log n) 100 × O(1)。当 n 很大时一次 O(n log n) 往往比 100 次 O(n) 更划算。具体怎么选可以按这个经验判断如果查询次数 m 和 log n 在同一量级或更大直接排序更划算。n1000 万时 log₂n 约 23也就是说查询次数超过大概 20 次排序的摊销成本就更优了。还有一种中间方案只做一次三路划分把数据切成若干桶后续查询落在哪个桶里就只处理那个桶。这本质上是把一次选择的结果复用起来适合查询次数中等、内存有限的场景。7.2 数据流中的第k大两个堆的配合真实业务里的第 k 小/大很多时候不是对静态数组查一次而是数据持续到达、需要动态维护。比如实时监控中维护最近一段时间的中位数数据每秒新增几百条。这时候快速选择就不适用了因为每次新数据到来都要重新跑一遍。标准解法是双堆一个大顶堆存较小的一半一个小顶堆存较大的一半维持两个堆的大小差不超过 1。插入新元素时先按值决定放进哪个堆然后调整大小平衡。这样中位数始终在两个堆顶之一或者两者的平均值上插入代价 O(log n)查询代价 O(1)。如果要维护的是第 k 大k 固定可以简化成维护一个大小为 k 的小顶堆堆顶就是第 k 大插入代价 O(log k)。这是我在日志系统里用得最多的结构统计最近 1000 条请求里耗时最高的 10 条用一个大小为 10 的小顶堆每条新日志进来和堆顶比较一下就行内存占用恒定在 k。7.3 内存受限时的分块思路还有一种情况值得单独提数据大到内存装不下。这时候前面所有在内存里做的分区都失效了需要把分治思想搬到外存上。思路是分两遍扫描第一遍把数据按值域切成若干块统计每块包含多少元素。比如值域是 0 到 10 亿切成 1000 个区间每个区间统计计数。这一步是纯流式的内存只需要 1000 个计数器。第二遍从 k 出发累加各区间的计数找到目标所在的区间然后只把这一个区间的数据读进来如果还是太大就再切一层在内存里用快速选择求第 k 小。这样总的 I/O 次数是常数级内存占用可控。这个方案我在处理一份几十 GB 的日志时用过原理不复杂但实现时要注意两点一是分块的值域边界要均匀避免某个块远大于其他块二是读第二遍时最好按块并行处理否则单线程读盘会成为瓶颈。最后分享一个我自己的体会快速选择这个算法代码量不到 40 行但它是少数几个复杂度优势能直接转化成业务收益的算法。我在一次数据统计模块的优化里把全排序换成快速选择单次统计从 1.2 秒降到 90 毫秒左右机器数量直接减了一半。但换来这个收益的前提是你得真的搞清楚它的退化条件、重复元素的处理方式还有边界怎么写。这三个地方任何一处出错收益归零还可能带来更难查的线上问题。所以我现在的习惯是任何用了分区思想的代码先写对拍脚本跑 2000 组随机用例再拿大量重复元素和近乎有序的数据各测一遍确认没有性能悬崖才算收工。
返回列表