You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何实现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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.20 03:10:36