题解:P17234 [Algo Beat Contest 017 C] 交互题

· · 题解

思路

考虑一个合法的区间 [l,r],其 \text{mex} 一定满足以下三点:

说人话,就是比 \text{mex} 小的,一定都存在,并且在合法区间内。\text{mex} 一定存在,但一定在合法区间外。正确性显然。

首先对于 \text{mex}=0 的情况,显然每两个 0 之间的区间的子区间都是合法的。简单统计即可。

然后对于 \text{mex}\neq 0 的情况,设数字 i 第一次出现和最后一次出现的位置分别为 l_i,r_i,考虑从小到大枚举 \text{mex},统计出当前合法区间必须包含的位置(即 [l_0,r_0][l_{\text{mex}-1},r_{\text{mex}-1}] 的并集,设为 A),接着考虑此区间与 [l_{\text{mex}},r_{\text{mex}}](设为 B)的关系:

对于这两种情况,即两个区间交叉或者 BA 的子集,显然无解。

对于这两种情况,即 AB 的交集为空集,合法区间的数量也是显然的。

这种情况最为复杂,我们需要先确认 A 区间内是否存在 \text{mex},如果存在,合法方案为 0。如果不存在,需要找到 [l_1,l_2] 中最靠右的 \text{mex}(位置记为 x),以及 [r_2,r_1] 中最靠左的 \text{mex}(位置记为 y)。这三部分均可以使用分块O(n\sqrt{n}) 的时空复杂度解决。

那么对于 \forall l\in (x,l_2],r\in [r2,y),区间 [l,r] 均是合法区间。

代码

#include<bits/stdc++.h>
using namespace std;
#define ull unsigned long long
#define ll long long
#define ld long double
#define dd double
#define pii pair<int,int>
#define pll pair<int,int>
//char buf[1<<23],*p1=buf,*p2=buf;
//#define getchar() (p1==p2&&(p2=(p1=buf)+fread(buf,1,1<<23,stdin),p1==p2)?EOF:*p1++)
inline int read() {
    int x = 0, f = 1;
    char ch;
    while (((ch = getchar()) < 48 || ch > 57) && ch != EOF)if (ch == '-')f = -1;
    if (ch == EOF)x = EOF;
    while (ch >= 48 && ch <= 57)x = x * 10 + ch - 48, ch = getchar();
    return x * f;
}
char __sta[1009], __len;
inline void write(ll x, int bo) {
    if (x < 0)putchar('-'), x = -x;
    do __sta[++__len] = x % 10 + 48, x /= 10;
    while (x);
    while (__len)putchar(__sta[__len--]);
    if (bo == 3)return;
    putchar(bo ? '\n' : ' ');
}
const int N=2e5+999,INF=2e9;
int n;
int a[N];
struct QUJIAN{
    int l=INF,r;
}num[N];
int maxnum;
//io
void input(){
    n=read();
    for(int i=1;i<=n;i++){
        a[i]=read();
        num[a[i]].l=min(num[a[i]].l,i);
        num[a[i]].r=max(num[a[i]].r,i);
    }
}
//init
ll sum[N];
void init(){
    for(int i=1;i<=n;i++){
        sum[i]=sum[i-1]+i;
    }
}
//分块
const int lgN=509,k=450;
bool ka[lgN][N];
int L(int id){
    return (id-1)*k+1;
}
int R(int id){
    return min(n,id*k);
}
int pos(int x){
    return (x-1)/k+1;
}
void build(){
    for(int i=1;i<=n;i++){
        ka[pos(i)][a[i]]=1;
    }
}
bool query_find(int l,int r,int v){//寻找此区间是否存在v
    if(pos(l)==pos(r)){
        for(int i=l;i<=r;i++){
            if(a[i]==v)return 1;
        }
        return 0;
    }else{
        for(int i=l;i<=R(pos(l));i++){
            if(a[i]==v)return 1;
        }
        for(int i=pos(l)+1;i<pos(r);i++){
            if(ka[i][v])return 1;
        }
        for(int i=L(pos(r));i<=r;i++){
            if(a[i]==v)return 1;
        }
        return 0;
    }
}
int query_l(int x,int v){//找到在[1,x]之间最靠右的v
    for(int i=x;i>=L(pos(x));i--){
        if(a[i]==v)return i;
    }
    for(int i=pos(x)-1;i>0;i--){
        if(ka[i][v]){
            for(int j=R(i);j>=L(i);j--){
                if(a[j]==v)return j;
            }
        }
    }
    return 0;
}
int query_r(int x,int v){//找到[x,n]之间最靠左的v
    for(int i=x;i<=R(pos(x));i++){
        if(a[i]==v)return i;
    }
    for(int i=pos(x)+1;i<=pos(n);i++){
        if(ka[i][v]){
            for(int j=L(i);j<=R(i);j++){
                if(a[j]==v)return j;
            }
        }
    }
    return n+1;
}
//solve
ll ans;
int nl=INF,nr;
void solve(){
    if(num[0].r==0){
        puts("0");
        return;
    }
    int las=0;
    for(int i=1;i<=n;i++){
        if(a[i]==0){
            ans+=sum[(i-1)-las];
            las=i;
        }
    }
    if(a[n]!=0){
        ans+=sum[n-las];
    }
    nl=num[0].l,nr=num[0].r;
    for(int mex=1;num[mex].r!=0;mex++){
        if(num[mex].r<nl){
            ans+=1LL*(nl-num[mex].r)*(n-nr+1);
        }else if(num[mex].l>nr){
            ans+=1LL*(num[mex].l-nr)*(nl);
//      }else if(!query_find(1,1,n,nl,nr,mex)){
//          int minn=query_l(1,1,n,nl,mex);
//          int maxx=query_r(1,1,n,nr,mex);
        }else if(!query_find(nl,nr,mex)){
            int minn=query_l(nl,mex);
            int maxx=query_r(nr,mex);
            ans+=1LL*(nl-minn)*(maxx-nr);
        }
        nl=min(nl,num[mex].l),nr=max(nr,num[mex].r);
//      cout<<ans<<endl;
    }
    write(ans,1);
}
int main() {
    input();
    init();
    build();
    solve();
    return 0;
}