如何实现Kiwiel改进的Floyd-Rivest中位数查找算法?
实现Kiwiel改进版Floyd-Rivest算法(迭代式)及优化建议
问题背景
Python标准库statistics.median()依赖全排序计算中位数,大数据集下效率受限。原生Floyd-Rivest算法平均时间复杂度为O(n),且能抵御Median of 3s Killer Sequence,但在处理含大量重复元素的数据集(如[1,1,1,1,2,3,4,5])时表现不佳。Krzysztof C. Kiwiel在《On Floyd and Rivest's SELECT algorithm》中提出改进方案,解决了重复数据的处理问题。
本文提供迭代式实现以避免递归栈溢出,并给出算法逻辑层面的提速与替代方案建议。
迭代式Kiwiel改进版Floyd-Rivest实现
Kiwiel的核心改进是在分区阶段引入双指针,分别标记等于基准值的左、右边界,避免原生算法在重复元素集中的无效分区操作。以下是迭代版本的实现:
from math import exp, log, sqrt from typing import Iterable, Sequence def sign(value: int | float) -> int: return bool(value > 0) - bool(value < 0) def swap(sequence: list[int | float], x: int, y: int) -> None: sequence[x], sequence[y] = sequence[y], sequence[x] def floyd_rivest_kiwiel(sequence: list[int | float], left: int, right: int, k: int) -> int | float: # 使用栈模拟递归,避免栈溢出 stack = [(left, right)] while stack: current_left, current_right = stack.pop() if current_right <= current_left: continue # 当区间过大时,先缩小范围 if current_right - current_left > 600: n = current_right - current_left + 1 i = k - current_left + 1 z = log(n) s = 0.5 * exp(2 * z / 3) sd = 0.5 * sqrt(z * s * (n - s) / n) * sign(i - n / 2) new_left = max(current_left, int(k - i * s / n + sd)) new_right = min(current_right, int(k + (n - i) * s / n + sd)) # 将原区间和新区间压栈,先处理原区间(栈后进先出) stack.append((current_left, current_right)) stack.append((new_left, new_right)) continue # Kiwiel改进的分区逻辑:处理重复元素 t = sequence[k] l, r = current_left, current_right p, q = current_left, current_right # 分区核心:将元素分为 <t, =t, >t 三部分 while True: while sequence[r] > t: r -= 1 while l <= r and sequence[l] <= t: if sequence[l] == t: swap(sequence, p, l) p += 1 l += 1 if l > r: break swap(sequence, l, r) r -= 1 # 将等于t的元素移动到中间区域 m = min(p - current_left, l - p) for i in range(m): swap(sequence, current_left + i, l - 1 - i) # 更新区间边界,确定k所在的分区 new_l = current_left + (l - p) new_r = l - 1 if k < new_l: stack.append((current_left, new_l - 1)) elif k > new_r: stack.append((new_r + 1, current_right)) else: # k在等于t的区间内,直接返回t return t return sequence[k] def median(data: Iterable[int | float] | Sequence[int | float]) -> int | float: sequence = list(data) length = len(sequence) if length == 0: raise ValueError("median requires at least one data point") end = length - 1 midpoint = end // 2 if length % 2 == 1: return floyd_rivest_kiwiel(sequence, 0, end, midpoint) else: # 偶数个元素时,取中间两个数的平均值 left_mid = floyd_rivest_kiwiel(sequence, 0, end, midpoint) right_mid = floyd_rivest_kiwiel(sequence, 0, end, midpoint + 1) return (left_mid + right_mid) / 2
算法逻辑层面的提速建议
- 动态调整区间缩小阈值:原生算法的600阈值可根据数据特征调整,比如含大量重复元素的数据集可将阈值调小(如200),提前进入高效分区阶段。
- 基准值优化:在小范围内(如3个随机元素)取中位数作为基准值,减少重复数据下基准值偏斜的概率,提升分区效率。
- 提前终止判断:在进入分区前,检查当前区间内所有元素是否相等(可通过首尾元素快速判断,再抽样验证),若相等直接返回该值,跳过后续操作。
- 分区指针优化:在移动指针时,直接跳过连续等于基准值的元素,减少不必要的交换操作。
替代算法建议
- Introselect算法:结合快速选择和中位数中位数算法,在平均情况下保持O(n)的效率,同时保证最坏情况下的O(n)时间复杂度(避免快速选择的最坏O(n²))。实现时可设置迭代次数阈值,当超过阈值时切换为中位数中位数选择。
- 中位数中位数选择:将数据集划分为每5个元素一组,取每组中位数,再递归取这些中位数的中位数作为基准值,保证最坏O(n)时间,适合对最坏复杂度有严格要求的场景,但常数因子略高于Floyd-Rivest。
内容的提问来源于stack exchange,提问作者VoidTwo
相关产品推荐
相关产品推荐

