题解 P1923 【深基9.例4】求第 k 小的数

· · 题解

提供一个最坏情况下也能 O(n) 的算法,该算法来自《算法导论》。

前言

我们知道,快速排序的时间复杂度是 O(n\log n),因为对应递推式是 T(n)=2T(\frac{n}{2})+O(n),如果每次只递归一半,递推式就变为 T(n)=2T(\frac{n}{2})+O(n),求出的复杂度就变为 O(n),这样就能利用快速排序的思路在平均 O(n) 的时间内求出第 k 大元素。

但是快速排序最坏情况下,也就是划分不均匀的情况下,时间复杂度会变成 O(n^2),这样求第 k 大元素也会变成 O(n^2) ,于是就有了下面一个算法,这个算法保证每次划分都是足够均匀的,从而保证了最坏情况下的时间复杂度。

算法描述

定义过程SELECT,它包括以下步骤:

  1. 将输入数组的 n 个元素划分为 \lfloor\frac{n}{5}\rfloor 组,每组 5 个元素,且至多只有一组由剩下的 n\ mod\ 5个元素组成。

  2. 寻找这 \lceil\frac{n}{5}\rceil 组中每一组的中位数:首先对每组元素进行插入排序,然后确定每组有序元素的中位数。

  3. 对第 2 步中找出的 \lceil\frac{n}{5}\rceil 个中位数,递归调用SELECT以找出其中位数 x(如果有偶数个中位数,为了方便,约定 x 是较小的中位数)。

  4. 利用修改过的PARTITION版本,按中位数的中位数 x 对输入数组进行划分。让 k 比划分的低区中的元素数目多 1,因此 x 是第 k 小的元素,并且有 n-k 个元素在划分的高区。

  5. 如果 i=k,则返回 x。如果 i<k,则在低区递归调用SELECT来找出第 i 小的元素。如果 i>k,则在高区递归调用SELECT来找出第 i-k 小的元素。

第四步的PARTITION过程是快速排序的划分过程,但是指定用来划分的元素。

算法正确性及复杂度

算法正确性很容易证明,这里略过。

关于复杂度,可以这样思考:利用每组的中位数,可以得知大约一半的数比选出的 x 小,因此划分是均匀的。

详细证明在这里:

在这 \lceil \frac{n}{5} \rceil 个组中,除了当 n 不能被 5 整除时产生的所含元素少于 5 的那个组和包含 x 的那个组之外,至少有一半的组中有 3 个元素大于 x。不算这两个组,大于 x 的元素个数至少为 3(\lceil\frac{1}{2}\lceil \frac{n}{5}\rceil\rceil-2)\geq \frac{3n}{10}-6。类似地,至少有 \frac{3n}{10}-6 个元素小于 x 。因此在最坏情况下,在第 5 步中SELECT的递归调用至多作用于 \frac{7n}{10}+6 个元素。

因此,可以得到递推式 T(n)\leq \begin{cases} O(1) & n<140 \\ T(\lceil\frac{n}{5}\rceil)+T(\frac{7n}{10}+6)+O(n) & n\geq 140 \end{cases}

可以证明这个递推式的解为 T(n)=O(n),证明过程略。

数字 140 是算法导论中选取的数,也可以用任何严格大于 70 的数。

算法细节

首先,递归的时候,如果 n 很小,就直接插入排序求答案,不要再递归。否则,在特定情况下,可能会出现无限递归。

另外,划分过程要划分为三部分(大于 x,等于 x,小于 x),否则也可能造成无限递归。

最后,这个算法常数较大(不吸氧用时820毫秒),因此并不很实用,但万一出题人照着代码造数据呢

代码

#include<cstdio>
int read(){
    int n=0;char c=getchar();
    while(c>='0'&&c<='9'){n=n*10+c-'0';c=getchar();}
    return n;
}
int a[5000000];
void insert_sort(int beg,int end){
    for(int i=beg+1;i<end;i++){
        int j=i;
        while(j!=beg && a[j]<a[j-1]){
            int t=a[j];a[j]=a[j-1];a[j-1]=t;j--;
        }
    }
}
const int GROUP_SIZE=5;
const int MAX_SORT_SIZE=140;
int find_kth(int beg,int end,int k){
    if(end-beg<=MAX_SORT_SIZE){
    //如果问题规模较小,就直接插入排序求解
        insert_sort(beg,end);
        return a[beg+k];
    }else{
        int len2=(end-beg+GROUP_SIZE-1);
        int pos=beg;
        //先把每组的中位数找出来
        for(int t=beg;t<end;t+=GROUP_SIZE){
            int e1=t+GROUP_SIZE;
            if(e1>end)e1=end;
            insert_sort(t,e1);
            int id=(e1-t-1)/2;
            int s=a[t+id];a[t+id]=a[pos];a[pos]=s;
            pos++;
        }
        //找中位数的中位数
        int val=find_kth(beg,pos,(beg+pos)/2-beg);
        //将序列划分为三部分
        int p1=beg,p2=end,p3=end;
        while(p1<p2){
            if(a[p1]<val){
                p1++;
            }else{
                p2--;
                int t=a[p1];a[p1]=a[p2];a[p2]=t;
            }
        }
        while(p2<p3){
            if(a[p2]==val)p2++;
            else{
                p3--;
                int t=a[p2];a[p2]=a[p3];a[p3]=t;
            }
        }
        //根据情况选择递归处理的部分
        if(beg+k>=p1){
            if(beg+k>=p2)return find_kth(p2,end,k-(p2-beg));
            else return val;
        }else{
            return find_kth(beg,p1,k);
        }
    }
}
int main(){
    int n=read();int k=read();
    for(int i=0;i<n;i++){
        a[i]=read();
    }
    printf("%d\n",find_kth(0,n,k));
}

附加内容

问题也是算法导论上的,这里不给出答案。

  1. 如果算法改为每组 7 个元素,该算法还是线性时间吗?证明:如果每组 3 个元素,那么算法不是线性的。

  2. 假设元素都是互异的,如何使快速排序在最坏情况下的时间复杂度为 O(n \log n)