Skip to content

最小空间中位数算法

img

这次的神秘算法是:知道输入数组大小 n,只能遍历一遍输入数组,如何在 n2+2\lfloor \frac{n}{2} \rfloor + 2 额外空间、线性时间复杂度的限制下求中位数。

参考论文是 A Minimal Space Selection Algorithm That Runs in Linear Time。这论文似乎没有 DOI 号,所以我不加论文链接了。

这篇论文是 1977 年的,当时还有磁带这个东西。只读、只能顺序遍历,倒带的代价比较大,这个算法看起来会合理一些。现在的硬盘可以随机访问,这个算法也基本只剩图一乐了。

1. 前置算法

1.1. 区间旋转

把两个相邻区间 [A B] 原地变成 [B A],保持区间内部顺序不变。经典的做法是三次翻转法(或手摇算法),这里就不展开了。

代码里直接调用标准库的 std::rotate,复杂度 O(n)O(n)

1.2. 选择算法

选择算法就是求第 k 小的数,在我之前的文章《O(n) 原地选择的不稳定版本》讲过 BFPRT 和消除递归栈版本的 BFPRT,这里可以当成黑盒调用。

代码里直接用了 std::nth_element 作为 stub 调用。

cpp
template <typename RandomIt, typename Proj = std::identity>
void select_stub(RandomIt first, RandomIt mid, RandomIt last, Proj proj = {}) {
    std::ranges::nth_element(first, mid, last, {}, proj);
}

2. 最小空间中位数算法

我们假设 n 是偶数时,n 个数里排名最中间的两个数都是中位数。

2.1. 引理

n 个数里任选 n2+1+i\lfloor \frac n 2 \rfloor + 1 + i 个数,它们的最小 i 个数和最大 i 个数一定不是中位数。

证明一下,首先 n 是奇数时,n 个数里比中位数小的数有 n2\lfloor \frac n 2 \rfloor 个。

如果任选数里第 k 大的数是中位数,在任选数里比中位数小的个数就是 n2+1+ik\lfloor \frac n 2 \rfloor + 1 + i - k,于是有:

n2+1+ikn2\lfloor \frac n 2 \rfloor + 1 + i - k \le \lfloor \frac n 2 \rfloor

化简:

ki+1k \ge i + 1

因此,如果任选数里有中位数,那么最大 i 个数一定不是中位数。对称可知,最小 i 个数也是一样。

img


n 是偶数时,n 个数里比上中位数小的数最多 n2\lfloor \frac n 2 \rfloor 个。

如果第 k 大的数是上中位数,在任选数里比它小的个数就是 n2+1+ik\lfloor \frac n 2 \rfloor + 1 + i - k,于是有:

n2+1+ikn2\lfloor \frac n 2 \rfloor + 1 + i - k \le \lfloor \frac n 2 \rfloor

化简:

ki+1k \ge i + 1

因此对上中位数可证,下中位数也是差不多的证明过程。

2.2. 算法思路

我们有 n2+2\lfloor \frac n 2 \rfloor + 2 大小的缓冲区,正好对应引理 i=1i = 1 的情况。

因此很快能想到一个初步思路,先读 n2+1\lfloor \frac n 2 \rfloor + 1 个数到缓冲区里。此时缓冲区还空 1 个位置,之后不断读 1 个数并淘汰 1 个最小值和 1 个最大值,最后就恰好剩下中位数。

img

但问题是,我们做不到维护一个结构,可以 O(1)O(1) 完成插入、查询或删除最小最大值。因为能做到就能完成 O(n)O(n) 的比较排序了(比较排序下界是 O(nlogn)O(n\log n))。

并且除了元素只能申请 O(1)O(1) 的额外内存,加了这个要求可能连 O(logn)O(\log n) 都很难做到。


我们发现,第 1 轮淘汰后,空出来的位置有 2 个。于是第二轮淘汰时,可以读 2 个数,淘汰 2 个最小值和 2 个最大值。这样空出来的位置会指数上升,第 k 轮淘汰后,缓冲区会空出 2k2^k 个位置。

大概证明就是由于第 k 轮之前淘汰了 21+22+...+2k1=2k22^1+2^2+...+2^{k-1}=2^k-2 个数,所以第 k 轮对应引理 i=2k1i=2^{k-1} 且 n 减少 2k22^k-2 的情况。

img


那么这有什么用呢?这就是算法最核心的地方了。

一开始读 n2+1\lfloor \frac n 2 \rfloor + 1 个数后,我们先求 1 个最小值 L1L_1、1 个最大值 U1U_1,剩下的数求 2 个最小值 L2L_2、2 个最大值 U2U_2,剩下的数求 4 个最小值 L3L_3、4 个最大值 U3U_3,按等比数列类推。

img

到了淘汰阶段,第一轮读 1 个数,它和 L1,U1L_1,U_1 放一起(3 个数),并淘汰这里的 1 个最小值 1 个最大值。显然缓冲区剩下的数都夹在 L_1、U_1 之间,不可能出现最小或最大值。

第一轮有 2 个数淘汰,1 个数晋级。

第二轮淘汰,读 2 个数,它们和晋级的那个数、L2,U2L_2,U_2 放一起(7 个数),并淘汰这里的 2 个最小值 2 个最大值。同样缓冲区剩下的数都夹在 L2,U2L_2,U_2 之间,不可能出现最小 2 个或最大 2 个值。

第二轮有 4 个数淘汰,3 个数晋级。

第 i 轮也是按等比数列类推。读 2i12^{i-1} 个数,它们和前一轮晋级的 2i112^{i-1}-1 个数、Li,UiL_i,U_i 放一起(一共 2i+112^{i+1}-1 个数),并淘汰这里的 2i12^{i-1} 个最小值 2i12^{i-1} 个最大值。

img


算法大致思路就是这样,对于边界的处理,我在后文实现的时候再说。

2.3. 复杂度分析

初始化求 L1,U1,L2,U2,...L_1, U_1, L_2, U_2, ...,我们需要反过来计算,就是从排名中间的大区间再到两边的小区间。最大的两个区间大小大约是 n2\lfloor \frac n 2 \rfloor,把它们处理好后,剩下的范围直接减半。求区间本身是两次线性时间的选择算法,因此复杂度是一个等比数列 O(n)+O(n2)+O(n4)+...=O(n)O(n) + O(\frac n 2) + O(\frac n 4) + ...=O(n)

淘汰阶段,第 i 轮要处理 2i+112^{i+1}-1 个数,淘汰本身就是两次线性时间的选择算法,所以复杂度也是个等比数列 O(1)+O(2)+O(4)+...+O(n)=O(n)O(1) + O(2) + O(4) + ... + O(n)=O(n)

最终复杂度 O(n)O(n)

3. 实现最小空间中位数算法

理论可行,实践开始。

3.1. 选择一个区间

辅助函数,选择算法执行两次把一个区间划分出来。

cpp
template <typename RandomIt, typename Proj = std::identity>
void select_range(RandomIt first, RandomIt left, RandomIt right, RandomIt last, Proj proj = {}) {
    if (left != last) {
        select_stub(first, left, last, proj);
    }
    if (right != last) {
        select_stub(left, right, last, proj);
    }
}

3.2. 中位数算法

还是辅助函数,如果所有候选者都在缓冲区了,直接 select_range 把中位数求出来。

cpp
template <typename RandomIt, typename Proj = std::identity>
std::array<std::iter_value_t<RandomIt>, 2> median(RandomIt first, RandomIt last, Proj proj = {}) {
    int64_t size = last - first;
    assert_or_throw(size > 0);
    if (size > 2) {
        select_range(first, first + ((size - 1) / 2), first + (size / 2) + 1, last, proj);
    }
    std::array<std::iter_value_t<RandomIt>, 2> result = {first[(size - 1) / 2], first[size / 2]};
    if (proj(result[0]) > proj(result[1])) {
        std::swap(result[0], result[1]);
    }
    return result;
}

3.3. 初始化阶段

先读 n / 2 + 1 个数。

cpp
for (int64_t i = 0; i < (size / 2) + 1; i++) {
    buffer_first[i + 1] = next();
}

每次 select_range 中间一个区间,把该区间旋转到后面。做完后,从左往右正好是 L1,U1L_1,U_1 混一起的区间,L2,U2L_2,U_2 混一起的区间,等等。

结尾会有不足 Li,UiL_i,U_i 大小的几个数,放最后面。

cpp
BufferIt left = buffer_first + 1;
BufferIt right = buffer_last;
for (int64_t ladder_size = max_ladder_size; ladder_size > 0; ladder_size /= 2) {
    int64_t n_keeps = (ladder_size - 1) * 2;
    assert_or_throw(n_keeps < right - left);
    select_range(left, left + (n_keeps / 2), right - (n_keeps / 2), right, proj);
    right = std::rotate(left + (n_keeps / 2), right - (n_keeps / 2), right);
}
assert_or_throw(left == right);

3.4. 淘汰阶段

每一轮都尽可能多读,就是一直读到输入数组的末尾,或者缓冲区空位置填满。此时如果读完了,可以直接用 median 函数得到结果,结束算法。

淘汰数量就是读取量的 2 倍,处理数量是读取量的 4 倍减 1,或者整个缓冲区。

cpp
BufferIt candidates = buffer_first + 1;
while (true) {
    int64_t n_loads = std::min(remain, candidates - buffer_first);
    for (int64_t i = 0; i < n_loads; i++) {
        candidates--;
        *candidates = next();
    }
    if (remain == 0) {
        break;
    }
    int64_t n_drops = n_loads * 2;
    int64_t n_candidates = std::min((n_loads * 4) - 1, buffer_last - candidates);
    if (n_candidates == buffer_last - candidates) {
        assert_or_throw(n_candidates >= n_drops + (size % 2 == 0 ? 2 : 1));
    }
    BufferIt left = candidates;
    BufferIt right = buffer_first + n_candidates;
    select_range(left, left + (n_drops / 2), right - (n_drops / 2), right, proj);
    candidates = std::rotate(left + (n_drops / 2), right - (n_drops / 2), right);
}
return median(candidates, buffer_last, proj);

4. 稳定化改造

稳定的中位数算法,计算结果等价于稳定排序后取中间元素。

很显然,上面介绍的算法是是不稳定的,那么能不能让它稳定呢?很不幸,不能。

大致讲一下我的直觉。因为在淘汰阶段,我们期望淘汰的最小值“相同数初始位置靠前”,因此 LiL_i 必须比后面数的“相同数初始位置靠前”。但是读进来的数比缓冲区任何数都“相同数初始位置靠后”,它们和 LiL_i 混在一起就无法区分了。

5. 完整代码

完整实现测试

6. 结尾

讲起来很绕,实现只用了大约 100 行,还是比较简洁的。

这类算法会研究最小空间,很多操作都和原地有关,我还是把它们收录进原地算法系列吧。

Powered by VitePress | Theme by Vdoing