题解:P9673 [ICPC 2022 Jinan R] Quick Sort

· · 题解

题意简述

给定一个排列。严格按照题面给出的快速排序程序执行,求整个排序过程中调用 swap 的次数。

解题思路

题面程序使用 Hoare 划分。对当前区间 [l,r],先保存中点元素作为主元 v,然后执行以下过程:

若逐个移动指针,快速排序退化时会重复扫描很长的区间,总时间可能达到 O(n^2)。真正发生的交换没有这么多,因此要让指针直接跳到下一个停止位置。

用线段树维护当前排列每个区间的最小值和最大值。

寻找左指针时,若某个完整区间的最大值小于 v,其中所有元素都满足 a_i<v,可以整体跳过。否则继续向下,并优先搜索左儿子,即可找到指定范围内第一个满足 a_i\ge v 的位置。

寻找右指针时,若某个完整区间的最小值大于 v,也可以整体跳过。优先搜索右儿子,便能找到最后一个满足 a_j\le v 的位置。每次交换后,只需在线段树中修改两个位置。

下面证明线段树跳跃与原程序完全等价。

一次划分开始时,主元的值已经被单独保存。即使原中点位置后来参与交换,本轮使用的主元也不会改变。

假设某轮循环开始前,两种实现的排列和指针位置相同。左侧查询只跳过最大值小于 v 的区间,原程序也会依次越过其中每个位置;查询返回的第一个未被跳过位置满足 a_i\ge v,恰好是原左指针的停止位置。右侧查询同理。若两指针交叉,两种实现返回同一个 j;否则交换的是同一对元素。更新两个叶子后,下一轮循环的归纳条件仍然成立。

因此,每次划分的返回位置、交换内容和交换次数都与题面程序相同。对子区间继续归纳,整个快速排序的模拟结果也完全一致。

还需估计总交换次数。

设一次划分返回 p,左右子区间长度分别为:

\begin{aligned} s & =p-l+1 \\ t & =r-p \end{aligned}

考虑本轮中的一次交换,交换位置为 i<j。交换后 a_i\le va_j\ge v,并且这两个位置不会再被本轮交换。之后左指针最迟在 j 停止,右指针最早也不会越过 i。所以最终返回位置满足:

i\le p<j

本轮不同交换使用的左端点互不相同,右端点也互不相同。每次交换都在左右子区间各占用一个位置,因此交换次数至多为:

\min(s,t)

把本轮每次交换分别计费给较短子区间中的一个不同位置。某个位置每被计费一次,它下一层所在的区间长度至多是当前的一半。因此每个位置至多被计费 \lceil\log_2n\rceil 次,总交换次数为 O(n\log n)

除了实际交换,每次非单点划分还会进行最后一轮指针查询,所以查询总次数仍为 O(n\log n)。一次查询或单点修改的时间复杂度为 O(\log n),总时间复杂度为 O(n\log^2 n),空间复杂度为 O(n)

代码中的线段树结点用一个数对保存区间最小值和最大值。query 的最后一个参数为 0 时从左向右寻找大于等于主元的位置,为 1 时从右向左寻找小于等于主元的位置。

快速排序本身使用显式栈。先压入右区间,再压入左区间,出栈顺序便与原递归程序一致,同时不会出现 O(n) 层递归导致的栈溢出。

参考代码

#include <bits/stdc++.h>
using namespace std;

using pii=pair<int,int>;
const int N=500005;
int n,ans;
int a[N],stk_l[N],stk_r[N];
struct SEG
{
    pii val[N<<2];
    void push_up(int u)
    {
        val[u].first=min(val[u*2].first,val[u*2+1].first);
        val[u].second=max(val[u*2].second,val[u*2+1].second);
    }
    void build(int u,int l,int r)
    {
        if(l==r){val[u]={a[l],a[l]};return;}
        int mid=l+r>>1;
        build(u*2,l,mid);
        build(u*2+1,mid+1,r);
        push_up(u);
    }
    void update(int u,int l,int r,int p)
    {
        if(l==r){val[u]={a[l],a[l]};return;}
        int mid=l+r>>1;
        if(p<=mid)update(u*2,l,mid,p);
        else update(u*2+1,mid+1,r,p);
        push_up(u);
    }
    int query(int u,int l,int r,int x,int y,int v,bool rev)
    {
        if(r<x||y<l)return 0;
        if(!rev&&val[u].second<v)return 0;
        if(rev&&val[u].first>v)return 0;
        if(l==r)return l;
        int mid=l+r>>1;
        if(!rev)
        {
            int res=query(u*2,l,mid,x,y,v,rev);
            return res?res:query(u*2+1,mid+1,r,x,y,v,rev);
        }
        int res=query(u*2+1,mid+1,r,x,y,v,rev);
        return res?res:query(u*2,l,mid,x,y,v,rev);
    }
}seg;
int partition(int l,int r)
{
    int val=a[l+r>>1],i=l-1,j=r+1;
    while(1)
    {
        i=seg.query(1,1,n,i+1,r,val,0);
        j=seg.query(1,1,n,l,j-1,val,1);
        if(i>=j)return j;
        swap(a[i],a[j]);
        seg.update(1,1,n,i);
        seg.update(1,1,n,j);
        ans++;
    }
}
void solve()
{
    cin>>n;
    for(int i=1;i<=n;i++)cin>>a[i];
    seg.build(1,1,n);
    ans=0;
    int top=1;
    stk_l[1]=1;
    stk_r[1]=n;
    while(top)
    {
        int l=stk_l[top],r=stk_r[top--];
        if(l>=r)continue;
        int p=partition(l,r);
        if(p+1<r)
        {
            stk_l[++top]=p+1;
            stk_r[top]=r;
        }
        if(l<p)
        {
            stk_l[++top]=l;
            stk_r[top]=p;
        }
    }
    cout<<ans<<'\n';
}
int main()
{
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    int T;
    cin>>T;
    while(T--)solve();
    return 0;
}