题解:P9292 [ROI 2018] Robomarathon

· · 题解

题意描述

给定长度为 n 的数组 a,对于 [1,n]\cap\mathbb Z 的非空子集 S(称为“发炮集合”),起跑时间 x_i=\min\limits_{j\in S}|i-j|,冲线时间 b_i=a_i+x_i,排名 rk_i=1+\sum\limits_{j=1}^n [b_j<b_i]

## 题解 我们要最小/最大化排名 $rk_i$,而不仅仅是最小/最大化时间 $b_i$,这需要我们让 $i$ 的时间比别人小/大得更多。这一思想会同时用在两个 Version 中。 ### $p=1

要让排名最小,显然是仅在 i 处发炮。因为如果在别处发炮,其他人的 b_j 就会变小,rk_i 一定不会变优。

考虑如何快速统计排名。|i-j| 是这样的图像:

要求出 b_j<b_i 的个数,用树状数组维护即可(注意离散化)。当已完成 rk_i 的计算,i 即将增加 1 时,所有 j\le ib_j 都会增加 1,而 j>ib_j 都会减少 1。我们把这两侧分成 2 棵树状数组维护,像线段树的 lazy tag 一样打上偏移标记,在后续修改或查询时都减去此标记即可。

p=2

这个比较难想,我的考场想法是,在最远的地方发炮最优。然而这并不一定正确。考虑这组数据:

5 2
2 1 1 1 2

i=3 时,在第 1 个位置发炮,b=[2,2,3,4,6]rk_3=3。然而我们可以同时在位置 1,5 发炮,b=[2,2,3,2,2]rk_3=5

这是什么原理?还记得刚才说的吗,要比别人大得更多,排名才能更靠后。当 x_i=\min\limits_{j\in S}|i-j| 从红色线变为绿色线,别人的时间更短了,i 的排名就更可能变大。

于是想让 ix_i 在所有人中最大,可以采用这样的构造:设 d=i-1d=n-i,并在 i 两侧空出半径为 d-1 的区间,其余位置发炮。也即:当 i\le\lfloor\frac{n}{2}\rfloor 时在 \{1\}\cup[2i-1,n]\{n\} 发炮,当 i>\lfloor\frac{n}{2}\rfloor 时在 [1,2i-n]\cup\{n\}\{1\} 发炮。

考虑如何统计。下面以 i\le\lfloor\frac{n}{2}\rfloor 为例。我们分 3 个树状数组维护,当 i 即将增加 1 时,[1,i] 的部分不动,[i+1,2i-1] 的部分加 22i 这一点加 1,其余不动。接着倒着做一遍,再考虑在两端发炮的情况,就做完了。

代码比较难写,我封装了一个带离线询问 + 离散化的树状数组,会好写很多。请注意下面代码中 query(x) 查询的是严格小于 x 的个数。

时间复杂度 O(n\log n)

#include<bits/stdc++.h>
using namespace std;
int n,p,a[400010],ans[400010];
struct bit{
    int n,lazy,cur;
    vector<int> a,id,ans;
    struct oper{
        int type,x,val;
    };
    vector<oper> op;
    bit():lazy(0){}
    void add(int x){
        lazy+=x;
    }
    void upd(int x,int val){
        op.push_back({1,x-lazy,val});
    }
    void query(int x){
        op.push_back({2,x-lazy-1,0});
    }
    void o_upd(int x,int val){
        for(x++;x<=n;x+=x&-x) a[x]+=val;
    }
    int o_query(int x){
        int res=0;
        for(x++;x;x-=x&-x) res+=a[x];
        return res;
    }
    void calc(){
        n=op.size();
        for(int i=0;i<n;i++) id.push_back(i);
        stable_sort(id.begin(),id.end(),[&](int x,int y){
            return op[x].x<op[y].x;
        });
        for(int i=0;i<n;i++) op[id[i]].x=i;
        for(int i=0;i<=n;i++) a.push_back(0);
        cur=0;
        for(oper it:op){
            if(it.type==1) o_upd(it.x,it.val);
            if(it.type==2) ans.push_back(o_query(it.x));
        }
    }
    int get(){
        if(ans.empty()) calc();
        return ans[cur++];
    }
};
int main(){
    scanf("%d%d",&n,&p);
    for(int i=1;i<=n;i++) scanf("%d",&a[i]);
    if(p==1){
        bit rt1,rt2;
        for(int i=1;i<=n;i++) rt2.upd(a[i]+i-1,1);
        for(int i=1;i<=n;i++){
            rt1.query(a[i]),rt2.query(a[i]);
            rt1.upd(a[i],1);
            rt2.upd(a[i],-1);
            rt1.add(1),rt2.add(-1);
        }
        for(int i=1;i<=n;i++) ans[i]=rt1.get()+rt2.get()+1;
    }else{
        bit rt1,rt2,rt3;
        rt2.upd(a[1],1);
        for(int i=2;i<=n;i++) rt3.upd(a[i],1);
        for(int i=1;i<=n/2;i++){
            rt1.query(a[i]+i-1),rt2.query(a[i]+i-1),rt3.query(a[i]+i-1);
            rt1.upd(a[i]+i-1,1);
            rt2.upd(a[i]+i-1,-1);
            if(2*i<=n) rt2.upd(a[2*i]-1,1),rt3.upd(a[2*i],-1);
            if(2*i+1<=n) rt2.upd(a[2*i+1]-2,1),rt3.upd(a[2*i+1],-1);
            rt2.add(2);
        }
        for(int i=1;i<=n/2;i++) ans[i]=rt1.get()+rt2.get()+rt3.get()+1;
        rt1=bit(),rt2=bit(),rt3=bit();
        rt2.upd(a[n],1);
        for(int i=1;i<n;i++) rt3.upd(a[i],1);
        for(int i=1;i<=(n+1)/2;i++){
            rt1.query(a[n-i+1]+i-1),rt2.query(a[n-i+1]+i-1),rt3.query(a[n-i+1]+i-1);
            rt1.upd(a[n-i+1]+i-1,1);
            rt2.upd(a[n-i+1]+i-1,-1);
            if(n-2*i+1>=1) rt2.upd(a[n-2*i+1]-1,1),rt3.upd(a[n-2*i+1],-1);
            if(n-2*i>=1) rt2.upd(a[n-2*i]-2,1),rt3.upd(a[n-2*i],-1);
            rt2.add(2);
        }
        for(int i=1;i<=(n+1)/2;i++) ans[n-i+1]=rt1.get()+rt2.get()+rt3.get()+1;
        rt1=bit();
        for(int i=1;i<=n;i++) rt1.upd(a[i]+i-1,1);
        for(int i=1;i<=n;i++) rt1.query(a[i]+i-1);
        for(int i=1;i<=n;i++) ans[i]=max(ans[i],rt1.get()+1);
        rt1=bit();
        for(int i=1;i<=n;i++) rt1.upd(a[i]+n-i,1);
        for(int i=1;i<=n;i++) rt1.query(a[i]+n-i);
        for(int i=1;i<=n;i++) ans[i]=max(ans[i],rt1.get()+1);
    }
    for(int i=1;i<=n;i++) printf("%d\n",ans[i]);
}